diff --git a/examples/recipes/openai_whisper-medium/cpu/cpu/automatic-speech-recognition_fp16_decoder_config.json b/examples/recipes/openai_whisper-medium/cpu/cpu/automatic-speech-recognition_fp16_decoder_config.json new file mode 100644 index 000000000..153159bcf --- /dev/null +++ b/examples/recipes/openai_whisper-medium/cpu/cpu/automatic-speech-recognition_fp16_decoder_config.json @@ -0,0 +1,909 @@ +{ + "export": { + "opset_version": 17, + "batch_size": 1, + "export_params": true, + "do_constant_folding": true, + "verbose": false, + "dynamo": false, + "enable_hierarchy_tags": true, + "clean_onnx": false, + "hierarchy_tag_format": "full", + "input_tensors": [ + { + "name": "decoder_input_ids", + "dtype": "int32", + "shape": [ + 1, + 1 + ], + "value_range": [ + 0, + 51865 + ] + }, + { + "name": "encoder_hidden_states", + "dtype": "float32", + "shape": [ + 1, + 1500, + 1024 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "decoder_attention_mask", + "dtype": "bool", + "shape": [ + 1, + 448 + ] + }, + { + "name": "cache_position", + "dtype": "int64", + "shape": [ + 1 + ] + }, + { + "name": "past_0_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_0_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_1_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_1_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_2_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_2_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_3_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_3_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_4_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_4_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_5_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_5_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_6_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_6_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_7_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_7_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_8_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_8_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_9_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_9_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_10_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_10_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_11_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_11_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_12_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_12_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_13_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_13_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_14_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_14_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_15_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_15_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_16_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_16_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_17_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_17_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_18_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_18_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_19_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_19_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_20_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_20_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_21_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_21_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_22_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_22_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_23_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_23_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + } + ], + "output_tensors": [ + { + "name": "logits" + }, + { + "name": "present_0_key" + }, + { + "name": "present_0_value" + }, + { + "name": "present_1_key" + }, + { + "name": "present_1_value" + }, + { + "name": "present_2_key" + }, + { + "name": "present_2_value" + }, + { + "name": "present_3_key" + }, + { + "name": "present_3_value" + }, + { + "name": "present_4_key" + }, + { + "name": "present_4_value" + }, + { + "name": "present_5_key" + }, + { + "name": "present_5_value" + }, + { + "name": "present_6_key" + }, + { + "name": "present_6_value" + }, + { + "name": "present_7_key" + }, + { + "name": "present_7_value" + }, + { + "name": "present_8_key" + }, + { + "name": "present_8_value" + }, + { + "name": "present_9_key" + }, + { + "name": "present_9_value" + }, + { + "name": "present_10_key" + }, + { + "name": "present_10_value" + }, + { + "name": "present_11_key" + }, + { + "name": "present_11_value" + }, + { + "name": "present_12_key" + }, + { + "name": "present_12_value" + }, + { + "name": "present_13_key" + }, + { + "name": "present_13_value" + }, + { + "name": "present_14_key" + }, + { + "name": "present_14_value" + }, + { + "name": "present_15_key" + }, + { + "name": "present_15_value" + }, + { + "name": "present_16_key" + }, + { + "name": "present_16_value" + }, + { + "name": "present_17_key" + }, + { + "name": "present_17_value" + }, + { + "name": "present_18_key" + }, + { + "name": "present_18_value" + }, + { + "name": "present_19_key" + }, + { + "name": "present_19_value" + }, + { + "name": "present_20_key" + }, + { + "name": "present_20_value" + }, + { + "name": "present_21_key" + }, + { + "name": "present_21_value" + }, + { + "name": "present_22_key" + }, + { + "name": "present_22_value" + }, + { + "name": "present_23_key" + }, + { + "name": "present_23_value" + } + ], + "compatibility": { + "transformers_attention": "eager" + } + }, + "optim": {}, + "quant": { + "mode": "fp16", + "samples": 10, + "calibration_method": "minmax", + "weight_type": "uint8", + "activation_type": "uint8", + "per_channel": false, + "symmetric": false, + "weight_symmetric": null, + "activation_symmetric": null, + "save_calibration": false, + "distribution": "uniform", + "seed": null, + "calibration_load_path": null, + "calibration_save_path": null, + "op_types_to_quantize": null, + "nodes_to_exclude": null, + "task": "text2text-generation", + "model_id": "openai/whisper-medium", + "model_type": "whisper", + "fp16_keep_io_types": true, + "fp16_op_block_list": null + }, + "compile": null, + "loader": { + "task": "text2text-generation", + "model_class": "WhisperDecoderWrapper", + "model_type": "whisper" + } +} \ No newline at end of file diff --git a/examples/recipes/openai_whisper-medium/cpu/cpu/automatic-speech-recognition_fp16_encoder_config.json b/examples/recipes/openai_whisper-medium/cpu/cpu/automatic-speech-recognition_fp16_encoder_config.json new file mode 100644 index 000000000..d221bbdde --- /dev/null +++ b/examples/recipes/openai_whisper-medium/cpu/cpu/automatic-speech-recognition_fp16_encoder_config.json @@ -0,0 +1,66 @@ +{ + "export": { + "opset_version": 17, + "batch_size": 1, + "export_params": true, + "do_constant_folding": true, + "verbose": false, + "dynamo": false, + "enable_hierarchy_tags": true, + "clean_onnx": false, + "hierarchy_tag_format": "full", + "input_tensors": [ + { + "name": "input_features", + "dtype": "float32", + "shape": [ + 1, + 80, + 3000 + ], + "value_range": [ + 0, + 1 + ] + } + ], + "output_tensors": [ + { + "name": "encoder_hidden_states" + } + ], + "compatibility": { + "transformers_attention": "eager" + } + }, + "optim": {}, + "quant": { + "mode": "fp16", + "samples": 10, + "calibration_method": "minmax", + "weight_type": "uint8", + "activation_type": "uint8", + "per_channel": false, + "symmetric": false, + "weight_symmetric": null, + "activation_symmetric": null, + "save_calibration": false, + "distribution": "uniform", + "seed": null, + "calibration_load_path": null, + "calibration_save_path": null, + "op_types_to_quantize": null, + "nodes_to_exclude": null, + "task": "feature-extraction", + "model_id": "openai/whisper-medium", + "model_type": "whisper", + "fp16_keep_io_types": true, + "fp16_op_block_list": null + }, + "compile": null, + "loader": { + "task": "feature-extraction", + "model_class": "WhisperEncoderWrapper", + "model_type": "whisper" + } +} \ No newline at end of file diff --git a/examples/recipes/openai_whisper-medium/cpu/cpu/automatic-speech-recognition_fp32_decoder_config.json b/examples/recipes/openai_whisper-medium/cpu/cpu/automatic-speech-recognition_fp32_decoder_config.json new file mode 100644 index 000000000..09f62447c --- /dev/null +++ b/examples/recipes/openai_whisper-medium/cpu/cpu/automatic-speech-recognition_fp32_decoder_config.json @@ -0,0 +1,887 @@ +{ + "export": { + "opset_version": 17, + "batch_size": 1, + "export_params": true, + "do_constant_folding": true, + "verbose": false, + "dynamo": false, + "enable_hierarchy_tags": true, + "clean_onnx": false, + "hierarchy_tag_format": "full", + "input_tensors": [ + { + "name": "decoder_input_ids", + "dtype": "int32", + "shape": [ + 1, + 1 + ], + "value_range": [ + 0, + 51865 + ] + }, + { + "name": "encoder_hidden_states", + "dtype": "float32", + "shape": [ + 1, + 1500, + 1024 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "decoder_attention_mask", + "dtype": "bool", + "shape": [ + 1, + 448 + ] + }, + { + "name": "cache_position", + "dtype": "int64", + "shape": [ + 1 + ] + }, + { + "name": "past_0_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_0_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_1_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_1_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_2_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_2_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_3_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_3_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_4_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_4_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_5_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_5_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_6_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_6_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_7_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_7_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_8_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_8_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_9_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_9_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_10_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_10_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_11_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_11_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_12_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_12_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_13_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_13_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_14_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_14_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_15_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_15_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_16_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_16_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_17_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_17_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_18_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_18_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_19_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_19_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_20_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_20_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_21_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_21_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_22_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_22_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_23_key", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + }, + { + "name": "past_23_value", + "dtype": "float32", + "shape": [ + 1, + 16, + 448, + 64 + ], + "value_range": [ + 0, + 1 + ] + } + ], + "output_tensors": [ + { + "name": "logits" + }, + { + "name": "present_0_key" + }, + { + "name": "present_0_value" + }, + { + "name": "present_1_key" + }, + { + "name": "present_1_value" + }, + { + "name": "present_2_key" + }, + { + "name": "present_2_value" + }, + { + "name": "present_3_key" + }, + { + "name": "present_3_value" + }, + { + "name": "present_4_key" + }, + { + "name": "present_4_value" + }, + { + "name": "present_5_key" + }, + { + "name": "present_5_value" + }, + { + "name": "present_6_key" + }, + { + "name": "present_6_value" + }, + { + "name": "present_7_key" + }, + { + "name": "present_7_value" + }, + { + "name": "present_8_key" + }, + { + "name": "present_8_value" + }, + { + "name": "present_9_key" + }, + { + "name": "present_9_value" + }, + { + "name": "present_10_key" + }, + { + "name": "present_10_value" + }, + { + "name": "present_11_key" + }, + { + "name": "present_11_value" + }, + { + "name": "present_12_key" + }, + { + "name": "present_12_value" + }, + { + "name": "present_13_key" + }, + { + "name": "present_13_value" + }, + { + "name": "present_14_key" + }, + { + "name": "present_14_value" + }, + { + "name": "present_15_key" + }, + { + "name": "present_15_value" + }, + { + "name": "present_16_key" + }, + { + "name": "present_16_value" + }, + { + "name": "present_17_key" + }, + { + "name": "present_17_value" + }, + { + "name": "present_18_key" + }, + { + "name": "present_18_value" + }, + { + "name": "present_19_key" + }, + { + "name": "present_19_value" + }, + { + "name": "present_20_key" + }, + { + "name": "present_20_value" + }, + { + "name": "present_21_key" + }, + { + "name": "present_21_value" + }, + { + "name": "present_22_key" + }, + { + "name": "present_22_value" + }, + { + "name": "present_23_key" + }, + { + "name": "present_23_value" + } + ], + "compatibility": { + "transformers_attention": "eager" + } + }, + "optim": {}, + "quant": null, + "compile": null, + "loader": { + "task": "text2text-generation", + "model_class": "WhisperDecoderWrapper", + "model_type": "whisper" + } +} \ No newline at end of file diff --git a/examples/recipes/openai_whisper-medium/cpu/cpu/automatic-speech-recognition_fp32_encoder_config.json b/examples/recipes/openai_whisper-medium/cpu/cpu/automatic-speech-recognition_fp32_encoder_config.json new file mode 100644 index 000000000..82e0ce047 --- /dev/null +++ b/examples/recipes/openai_whisper-medium/cpu/cpu/automatic-speech-recognition_fp32_encoder_config.json @@ -0,0 +1,44 @@ +{ + "export": { + "opset_version": 17, + "batch_size": 1, + "export_params": true, + "do_constant_folding": true, + "verbose": false, + "dynamo": false, + "enable_hierarchy_tags": true, + "clean_onnx": false, + "hierarchy_tag_format": "full", + "input_tensors": [ + { + "name": "input_features", + "dtype": "float32", + "shape": [ + 1, + 80, + 3000 + ], + "value_range": [ + 0, + 1 + ] + } + ], + "output_tensors": [ + { + "name": "encoder_hidden_states" + } + ], + "compatibility": { + "transformers_attention": "eager" + } + }, + "optim": {}, + "quant": null, + "compile": null, + "loader": { + "task": "feature-extraction", + "model_class": "WhisperEncoderWrapper", + "model_type": "whisper" + } +} \ No newline at end of file diff --git a/src/winml/modelkit/commands/build.py b/src/winml/modelkit/commands/build.py index a56d9f52e..c17946189 100644 --- a/src/winml/modelkit/commands/build.py +++ b/src/winml/modelkit/commands/build.py @@ -2068,6 +2068,7 @@ def _build_hf_pipeline( max_iters: int = extra_kwargs.pop("hack_max_optim_iterations", 3) allow_unsupported_nodes: bool = extra_kwargs.pop("allow_unsupported_nodes", False) + skip_optimize: bool = extra_kwargs.pop("skip_optimize", False) model_label = model_id or "random-init" # ── Validate + setup ───────────────────────────────────────── @@ -2153,6 +2154,7 @@ def _name(base: str) -> str: show_io_first=False, analyze_output_path=analyze_result_path, allow_unsupported_nodes=allow_unsupported_nodes, + skip_optimize=skip_optimize or config.skip_optimize, ) # Persist config after autoconf @@ -2209,6 +2211,7 @@ def _build_onnx_pipeline( max_iters: int = extra_kwargs.pop("hack_max_optim_iterations", 3) allow_unsupported_nodes: bool = extra_kwargs.pop("allow_unsupported_nodes", False) + skip_optimize: bool = extra_kwargs.pop("skip_optimize", False) # ── Validate + setup ───────────────────────────────────────── if not onnx_path.exists(): @@ -2264,7 +2267,7 @@ def _build_onnx_pipeline( show_io_first=True, analyze_output_path=analyze_result_path, allow_unsupported_nodes=allow_unsupported_nodes, - skip_optimize=config.skip_optimize, + skip_optimize=skip_optimize or config.skip_optimize, ) config_path.write_text(json.dumps(config.to_dict(), indent=2)) diff --git a/src/winml/modelkit/eval/__init__.py b/src/winml/modelkit/eval/__init__.py index 435601a15..9894ca7d5 100644 --- a/src/winml/modelkit/eval/__init__.py +++ b/src/winml/modelkit/eval/__init__.py @@ -19,6 +19,9 @@ if TYPE_CHECKING: + from .automatic_speech_recognition_evaluator import ( + WinMLAutomaticSpeechRecognitionEvaluator, + ) from .depth_estimation_evaluator import WinMLDepthEstimationEvaluator from .feature_extraction_evaluator import WinMLFeatureExtractionEvaluator from .fill_mask_evaluator import WinMLFillMaskEvaluator @@ -47,6 +50,9 @@ _LAZY_ATTRS: dict[str, str] = { # Evaluators + "WinMLAutomaticSpeechRecognitionEvaluator": ( + ".automatic_speech_recognition_evaluator:WinMLAutomaticSpeechRecognitionEvaluator" + ), "WinMLDepthEstimationEvaluator": ".depth_estimation_evaluator:WinMLDepthEstimationEvaluator", "WinMLFeatureExtractionEvaluator": ( ".feature_extraction_evaluator:WinMLFeatureExtractionEvaluator" @@ -126,6 +132,7 @@ def __dir__() -> list[str]: "SpearmanCorrelationMetric", "TensorSimilarityEvaluator", "TopKAccuracyMetric", + "WinMLAutomaticSpeechRecognitionEvaluator", "WinMLDepthEstimationEvaluator", "WinMLEvaluationConfig", "WinMLEvaluator", diff --git a/src/winml/modelkit/eval/automatic_speech_recognition_evaluator.py b/src/winml/modelkit/eval/automatic_speech_recognition_evaluator.py new file mode 100644 index 000000000..040683633 --- /dev/null +++ b/src/winml/modelkit/eval/automatic_speech_recognition_evaluator.py @@ -0,0 +1,246 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- +"""Automatic speech recognition evaluation for CTC and seq2seq models.""" + +from __future__ import annotations + +import io +import math +from copy import deepcopy +from pathlib import Path +from typing import TYPE_CHECKING, Any, Literal, cast + +import numpy as np + +from ..utils.eval_utils import DatasetValidationError, get_default +from .base_evaluator import WinMLEvaluator + + +if TYPE_CHECKING: + from datasets import Dataset + from transformers.pipelines.base import Pipeline + + from .config import DatasetConfig, WinMLEvaluationConfig + + +ASRMode = Literal["ctc", "seq2seq"] + + +def _word_error_counts(prediction: str, reference: str) -> tuple[int, int]: + """Return word-level Levenshtein distance and reference word count.""" + predicted_words = prediction.lower().split() + reference_words = reference.lower().split() + if not reference_words: + return (0 if not predicted_words else len(predicted_words), 0) + + previous = list(range(len(predicted_words) + 1)) + for ref_index, ref_word in enumerate(reference_words, start=1): + current = [ref_index] + for pred_index, pred_word in enumerate(predicted_words, start=1): + current.append( + min( + current[-1] + 1, + previous[pred_index] + 1, + previous[pred_index - 1] + (ref_word != pred_word), + ) + ) + previous = current + return previous[-1], len(reference_words) + + +def _word_error_rate(prediction: str, reference: str) -> float: + """Return word-level Levenshtein distance divided by reference words.""" + errors, reference_words = _word_error_counts(prediction, reference) + if reference_words == 0: + return 0.0 if errors == 0 else 1.0 + return errors / reference_words + + +def _asr_mode(model: Any) -> ASRMode: + """Select ASR decoding from the model's architecture contract.""" + config = getattr(model, "config", None) + return "seq2seq" if bool(getattr(config, "is_encoder_decoder", False)) else "ctc" + + +class WinMLAutomaticSpeechRecognitionEvaluator(WinMLEvaluator): + """Evaluate raw speech with explicit CTC or encoder-decoder decoding.""" + + def __init__(self, config: WinMLEvaluationConfig, model: Any) -> None: + if config.model_id is None: + raise ValueError("ASR evaluation requires model_id to load its processor.") + + columns = config.dataset.columns_mapping + self._audio_column = columns.get( + "input_column", + cast("str", get_default("automatic-speech-recognition", "input_column")), + ) + self._label_column = columns.get( + "label_column", + cast("str", get_default("automatic-speech-recognition", "label_column")), + ) + self._max_audio_seconds = min( + 30.0, + float(columns.get("max_audio_seconds", "30")), + ) + self._max_new_tokens = min(32, int(columns.get("max_new_tokens", "32"))) + self._language = columns.get("language", "english") + self._generation_task = columns.get("generation_task", "transcribe") + if self._max_audio_seconds <= 0 or self._max_new_tokens <= 0: + raise ValueError("ASR audio and generation caps must be positive.") + + from transformers import AutoProcessor + + self.processor = AutoProcessor.from_pretrained( + config.model_id, + trust_remote_code=config.trust_remote_code, + ) + self.mode = _asr_mode(model) + super().__init__(config, model) + + def prepare_data(self) -> Dataset: + """Load rows with audio decoding disabled so bytes/path handling is explicit.""" + data = super().prepare_data() + try: + from datasets import Audio + + return cast("Dataset", data.cast_column(self._audio_column, Audio(decode=False))) + except (KeyError, TypeError, ValueError) as error: + raise DatasetValidationError( + f"could not expose raw audio column {self._audio_column!r}: {error}" + ) from error + + def prepare_pipeline(self) -> Pipeline | None: # type: ignore[override] + """ASR uses the processor and model directly to keep decoding paths explicit.""" + return None + + def align_labels(self, dataset: Dataset, ds_config: DatasetConfig) -> Dataset: + """Free-text transcripts need no class-label alignment.""" + return dataset + + def _decode_audio(self, value: Any) -> np.ndarray: + """Decode an Audio row, mix to mono, resample to 16 kHz, and cap duration.""" + sampling_rate: int | None = None + waveform: np.ndarray | None = None + + if isinstance(value, dict) and value.get("array") is not None: + waveform = np.asarray(value["array"], dtype=np.float32) + sampling_rate = int(value.get("sampling_rate") or 0) + else: + source: Any = value + if isinstance(value, dict): + if value.get("bytes") is not None: + source = io.BytesIO(value["bytes"]) + elif value.get("path"): + source = Path(value["path"]) + try: + import soundfile as sf + except ImportError as error: + raise DatasetValidationError( + "Raw ASR audio decoding requires the 'audio' extra: " + "pip install winml-cli[audio]" + ) from error + try: + decoded, sampling_rate = sf.read(source, dtype="float32", always_2d=False) + except Exception as error: + raise DatasetValidationError(f"failed to decode audio: {error}") from error + waveform = np.asarray(decoded, dtype=np.float32) + + if sampling_rate is None or sampling_rate <= 0 or waveform.size == 0: + raise DatasetValidationError("audio must contain samples and a positive sampling rate") + if waveform.ndim == 2: + waveform = waveform.mean(axis=1, dtype=np.float32) + if waveform.ndim != 1: + raise DatasetValidationError( + f"audio must be mono or channel-last, got {waveform.shape}" + ) + + target_rate = 16_000 + if sampling_rate != target_rate: + from scipy.signal import resample_poly + + divisor = math.gcd(sampling_rate, target_rate) + waveform = resample_poly( + waveform, + target_rate // divisor, + sampling_rate // divisor, + ).astype(np.float32, copy=False) + max_samples = int(target_rate * self._max_audio_seconds) + return np.ascontiguousarray(waveform[:max_samples], dtype=np.float32) + + def _predict_ctc(self, waveform: np.ndarray) -> str: + import torch + + encoded = self.processor(waveform, sampling_rate=16_000, return_tensors="pt") + outputs = self.model(**dict(encoded)) + logits = getattr(outputs, "logits", None) + if logits is None and isinstance(outputs, dict): + logits = outputs.get("logits") + if logits is None: + raise ValueError("CTC ASR model must return frame logits.") + token_ids = torch.as_tensor(logits).argmax(dim=-1) + decoded = self.processor.batch_decode(token_ids) + return str(decoded[0]).strip() if decoded else "" + + def _predict_seq2seq(self, waveform: np.ndarray) -> str: + encoded = self.processor(waveform, sampling_rate=16_000, return_tensors="pt") + generation_config = deepcopy(self.model.generation_config) + generation_config.max_new_tokens = self._max_new_tokens + generation_config.num_beams = 1 + generation_config.do_sample = False + if hasattr(self.processor, "get_decoder_prompt_ids"): + generation_config.language = self._language + generation_config.task = self._generation_task + generation_config.return_timestamps = False + generated = self.model.generate(**dict(encoded), generation_config=generation_config) + decoded = self.processor.batch_decode(generated, skip_special_tokens=True) + return str(decoded[0]).strip() if decoded else "" + + def compute(self) -> dict[str, Any]: + """Transcribe every selected row and report WER plus exact accounting.""" + predictions: list[str] = [] + references: list[str] = [] + for row_index, sample in enumerate(self.data): + audio = sample.get(self._audio_column) + reference = sample.get(self._label_column) + if audio is None or not isinstance(reference, str) or not reference.strip(): + raise DatasetValidationError( + f"row {row_index} must contain audio and a non-empty transcript" + ) + waveform = self._decode_audio(audio) + prediction = ( + self._predict_seq2seq(waveform) + if self.mode == "seq2seq" + else self._predict_ctc(waveform) + ) + predictions.append(prediction) + references.append(reference.strip()) + + if not predictions: + raise DatasetValidationError("ASR evaluation selected no usable rows") + word_counts = ( + _word_error_counts(prediction, reference) + for prediction, reference in zip(predictions, references, strict=True) + ) + error_count = 0 + reference_word_count = 0 + for errors, words in word_counts: + error_count += errors + reference_word_count += words + if reference_word_count == 0: + raise DatasetValidationError("ASR references must contain at least one word") + return { + "wer": error_count / reference_word_count, + "asr_mode": self.mode, + "requested_samples": self.config.dataset.samples, + "processed_samples": len(predictions), + "skipped_samples": 0, + "max_audio_seconds": self._max_audio_seconds, + "max_new_tokens": self._max_new_tokens if self.mode == "seq2seq" else 0, + "num_beams": 1 if self.mode == "seq2seq" else 0, + "return_sequences": 1, + } + + +__all__ = ["WinMLAutomaticSpeechRecognitionEvaluator"] diff --git a/src/winml/modelkit/eval/evaluate.py b/src/winml/modelkit/eval/evaluate.py index 6e5b4f223..bd4b08f2a 100644 --- a/src/winml/modelkit/eval/evaluate.py +++ b/src/winml/modelkit/eval/evaluate.py @@ -60,6 +60,8 @@ def _select_model_loader(config: WinMLEvaluationConfig) -> _ModelLoaderKind: # default formatter layout) yields >100-char lines that trip E501. # fmt: off _EVALUATOR_REGISTRY: dict[str, str] = { + "automatic-speech-recognition": + "winml.modelkit.eval.automatic_speech_recognition_evaluator:WinMLAutomaticSpeechRecognitionEvaluator", "image-classification": "winml.modelkit.eval.base_evaluator:WinMLEvaluator", "text-classification": @@ -177,6 +179,21 @@ def _validate_pytorch_runtime_config(config: WinMLEvaluationConfig) -> None: } _DEFAULT_DATASETS: dict[str, dict] = { + "automatic-speech-recognition": { + "path": "openslr/librispeech_asr", + "name": "clean", + "split": "validation", + "revision": "71cacbfb7e2354c4226d01e70d77d5fca3d04ba1", + "shuffle": False, + "columns_mapping": { + "input_column": "audio", + "label_column": "text", + "max_audio_seconds": "30", + "max_new_tokens": "32", + "language": "english", + "generation_task": "transcribe", + }, + }, "image-classification": { "path": "timm/mini-imagenet", "split": "test", diff --git a/src/winml/modelkit/loader/task.py b/src/winml/modelkit/loader/task.py index 0ae0ca6be..3de9ccb67 100644 --- a/src/winml/modelkit/loader/task.py +++ b/src/winml/modelkit/loader/task.py @@ -129,6 +129,7 @@ COMPOSITE_TASKS: frozenset[str] = frozenset( { "image-to-text", + "automatic-speech-recognition", "summarization", "table-question-answering", "text-generation", diff --git a/src/winml/modelkit/models/hf/__init__.py b/src/winml/modelkit/models/hf/__init__.py index 0128f05ab..a3e817b55 100644 --- a/src/winml/modelkit/models/hf/__init__.py +++ b/src/winml/modelkit/models/hf/__init__.py @@ -112,6 +112,9 @@ # triggers registration Wav2Vec2EmotionRegressionIOConfig as _Wav2Vec2EmotionRegressionIOConfig, ) +from .whisper import MODEL_CLASS_MAPPING as _WHISPER_CLASS_MAPPING +from .whisper import WhisperDecoderIOConfig as _WhisperDecoderIOConfig +from .whisper import WhisperEncoderIOConfig as _WhisperEncoderIOConfig from .zoedepth import ZoeDepthIOConfig as _ZoeDepthIOConfig # triggers registration @@ -148,6 +151,7 @@ _VED_CLASS_MAPPING, _VITPOSE_CLASS_MAPPING, _WAV2VEC2_CLASS_MAPPING, + _WHISPER_CLASS_MAPPING, ) for _key, _model_cls in _sub_mapping.items() } diff --git a/src/winml/modelkit/models/hf/whisper.py b/src/winml/modelkit/models/hf/whisper.py new file mode 100644 index 000000000..a7691bd10 --- /dev/null +++ b/src/winml/modelkit/models/hf/whisper.py @@ -0,0 +1,329 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- +"""Whisper split encoder/decoder export and inference support.""" + +from __future__ import annotations + +import logging +from types import SimpleNamespace +from typing import TYPE_CHECKING, Any, ClassVar, cast + +import torch +import torch.nn as nn +from optimum.exporters.onnx import OnnxConfig +from optimum.utils import NormalizedConfig +from optimum.utils.input_generators import DummyInputGenerator +from transformers import WhisperForConditionalGeneration +from transformers.cache_utils import DynamicCache, EncoderDecoderCache +from transformers.models.whisper.generation_whisper import WhisperGenerationMixin + +from ...export import register_onnx_overwrite +from ..winml.composite_model import register_composite_model +from ..winml.encoder_decoder import EncoderDecoderInputGenerator, WinMLEncoderDecoderCore +from ..winml.kv_cache import PastKeyValueInputGenerator, WinMLStaticCache + + +if TYPE_CHECKING: + from transformers import GenerationConfig, PretrainedConfig + from transformers.modeling_outputs import BaseModelOutput, Seq2SeqLMOutput + +logger = logging.getLogger(__name__) + + +class WhisperEncoderWrapper(nn.Module): + """Expose the Whisper audio encoder as a standalone model.""" + + def __init__(self, encoder: nn.Module) -> None: + super().__init__() + self.encoder = encoder + + @classmethod + def from_pretrained(cls, model_name_or_path: str, **kwargs: Any) -> WhisperEncoderWrapper: + """Load a conditional-generation checkpoint and retain its encoder.""" + full_model = WhisperForConditionalGeneration.from_pretrained(model_name_or_path, **kwargs) + wrapper = cls(full_model.get_encoder()) + wrapper.eval() + return wrapper + + def forward(self, input_features: torch.Tensor) -> torch.Tensor: + """Return the encoded audio hidden states.""" + return cast( + "torch.Tensor", + self.encoder(input_features=input_features).last_hidden_state, + ) + + +class WhisperDecoderWrapper(nn.Module): + """Expose Whisper decoding with WinML's static self-attention KV cache.""" + + def __init__(self, model: nn.Module, num_layers: int) -> None: + super().__init__() + self.model = model + self.num_layers = num_layers + self.config: PretrainedConfig = cast("PretrainedConfig", model.config) + + @classmethod + def from_pretrained(cls, model_name_or_path: str, **kwargs: Any) -> WhisperDecoderWrapper: + """Load a conditional-generation checkpoint for decoder export.""" + full_model = WhisperForConditionalGeneration.from_pretrained(model_name_or_path, **kwargs) + wrapper = cls(full_model, full_model.config.decoder_layers) + wrapper.eval() + return wrapper + + def get_export_args(self, inputs: dict[str, torch.Tensor]) -> tuple[torch.Tensor, ...]: + """Convert named dummy inputs to the positional export protocol.""" + return tuple(inputs.values()) + + def forward(self, *args: torch.Tensor) -> tuple[torch.Tensor, ...]: + """Decode one token and return logits plus new-token self-attention KV.""" + decoder_input_ids = args[0] + encoder_hidden_states = args[1] + decoder_attention_mask = args[2] + cache_position = args[3] + kv_start = 4 + + max_cache_len = args[kv_start].size(2) + self_attn_cache = WinMLStaticCache(self.config, max_cache_len=max_cache_len) + self_attn_cache.early_initialization( + batch_size=decoder_input_ids.size(0), + num_heads=args[kv_start].size(1), + head_dim=args[kv_start].size(3), + dtype=args[kv_start].dtype, + device=decoder_input_ids.device, + ) + for layer_index in range(self.num_layers): + layer = self_attn_cache.layers[layer_index] + layer.keys = args[kv_start + layer_index * 2] + layer.values = args[kv_start + layer_index * 2 + 1] + self_attn_cache.set_trace_position(cache_position) + + cache = EncoderDecoderCache(self_attn_cache, DynamicCache()) + output = self.model( + decoder_input_ids=decoder_input_ids, + encoder_outputs=(encoder_hidden_states,), + decoder_attention_mask=decoder_attention_mask, + past_key_values=cache, + use_cache=True, + cache_position=cache_position, + ) + + result: list[torch.Tensor] = [output.logits] + for layer_index in range(self.num_layers): + key, value = self_attn_cache.captured[layer_index] + result.extend((key, value)) + return tuple(result) + + +class WhisperEncoderInputGenerator(DummyInputGenerator): # type: ignore[misc] + """Generate Whisper log-Mel features at the architecture's fixed length.""" + + SUPPORTED_INPUT_NAMES = ("input_features",) + + def __init__( + self, + task: str, + normalized_config: NormalizedConfig, + batch_size: int = 1, + **kwargs: Any, + ) -> None: + self.batch_size = batch_size + self.feature_size = normalized_config.feature_size + self.audio_sequence_length = normalized_config.audio_sequence_length + + def generate( + self, + input_name: str, + framework: str = "pt", + int_dtype: str = "int64", + float_dtype: str = "fp32", + ) -> torch.Tensor: + """Generate the fixed-size feature tensor required by WhisperEncoder.""" + if input_name != "input_features": + raise ValueError(f"Unknown input: {input_name}") + return cast( + "torch.Tensor", + self.random_float_tensor( + (self.batch_size, self.feature_size, self.audio_sequence_length), + framework=framework, + dtype=float_dtype, + ), + ) + + +class _WhisperEncoderNormalizedConfig(NormalizedConfig): # type: ignore[misc] + FEATURE_SIZE = "num_mel_bins" + MAX_SOURCE_POSITIONS = "max_source_positions" + + @property + def audio_sequence_length(self) -> int: + return cast("int", self.max_source_positions * 2) + + +class _WhisperDecoderNormalizedConfig(NormalizedConfig): # type: ignore[misc] + VOCAB_SIZE = "vocab_size" + HIDDEN_SIZE = "d_model" + NUM_LAYERS = "decoder_layers" + NUM_ATTENTION_HEADS = "decoder_attention_heads" + MAX_CACHE_LEN = "max_target_positions" + ENCODER_SEQUENCE_LENGTH = "max_source_positions" + + @property + def head_dim(self) -> int: + return cast("int", self.hidden_size // self.num_attention_heads) + + @property + def sequence_length(self) -> int: + return cast("int", self.encoder_sequence_length) + + +@register_onnx_overwrite("whisper", "feature-extraction", library_name="transformers") +class WhisperEncoderIOConfig(OnnxConfig): # type: ignore[misc] + """ONNX configuration for the fixed-shape Whisper audio encoder.""" + + NORMALIZED_CONFIG_CLASS = _WhisperEncoderNormalizedConfig + DUMMY_INPUT_GENERATOR_CLASSES = (WhisperEncoderInputGenerator,) + + @property + def inputs(self) -> dict[str, dict[int, str]]: # noqa: D102 + return {"input_features": {0: "batch_size"}} + + @property + def outputs(self) -> dict[str, dict[int, str]]: # noqa: D102 + return {"encoder_hidden_states": {0: "batch_size"}} + + +@register_onnx_overwrite("whisper", "text2text-generation", library_name="transformers") +class WhisperDecoderIOConfig(OnnxConfig): # type: ignore[misc] + """ONNX configuration for token-at-a-time Whisper decoding.""" + + NORMALIZED_CONFIG_CLASS = _WhisperDecoderNormalizedConfig + DUMMY_INPUT_GENERATOR_CLASSES = ( + EncoderDecoderInputGenerator, + PastKeyValueInputGenerator, + ) + + @property + def inputs(self) -> dict[str, dict[int, str]]: # noqa: D102 + result: dict[str, dict[int, str]] = { + "decoder_input_ids": {0: "batch_size"}, + "encoder_hidden_states": {0: "batch_size"}, + "decoder_attention_mask": {0: "batch_size"}, + "cache_position": {}, + } + for layer_index in range(self._normalized_config.num_layers): + result[f"past_{layer_index}_key"] = {0: "batch_size"} + result[f"past_{layer_index}_value"] = {0: "batch_size"} + return result + + @property + def outputs(self) -> dict[str, dict[int, str]]: # noqa: D102 + result: dict[str, dict[int, str]] = {"logits": {0: "batch_size"}} + for layer_index in range(self._normalized_config.num_layers): + result[f"present_{layer_index}_key"] = {0: "batch_size"} + result[f"present_{layer_index}_value"] = {0: "batch_size"} + return result + + +MODEL_CLASS_MAPPING: dict[tuple[str, str], type] = { + ("whisper", "feature-extraction"): WhisperEncoderWrapper, + ("whisper", "text2text-generation"): WhisperDecoderWrapper, +} + + +@register_composite_model("whisper", "automatic-speech-recognition") +class WinMLWhisperModel( # type: ignore[misc] + WinMLEncoderDecoderCore, WhisperGenerationMixin +): + """Split Whisper transcription model backed by encoder and decoder ONNX graphs.""" + + main_input_name = "input_features" + _SUB_MODEL_CONFIG: ClassVar[dict[str, str]] = { + "encoder": "feature-extraction", + "decoder": "text2text-generation", + } + + def __init__( + self, + sub_models: dict[str, Any], + config: PretrainedConfig, + device: str = "cpu", + ) -> None: + super().__init__(sub_models, config, device) + encoder_frames = self._encoder._expected["input_features"][-1] + input_stride = encoder_frames // config.max_source_positions + self.model = SimpleNamespace( + config=config, + encoder=SimpleNamespace( + conv1=SimpleNamespace(stride=(1,)), + conv2=SimpleNamespace(stride=(input_stride,)), + ), + ) + model_name_or_path = getattr(config, "_name_or_path", "") + if model_name_or_path: + from transformers import GenerationConfig + + try: + self.generation_config = GenerationConfig.from_pretrained(model_name_or_path) + except OSError: + logger.warning( + "Could not load Whisper generation_config.json from %s; " + "language/task controls may be unavailable.", + model_name_or_path, + ) + + @classmethod + def get_cache_class(cls) -> type: # noqa: D102 + return WinMLStaticCache + + def forward( + self, + *, + input_features: torch.Tensor | None = None, + encoder_outputs: BaseModelOutput | tuple | None = None, + **kwargs: Any, + ) -> Seq2SeqLMOutput: + """Run the audio encoder when callers have not precomputed its outputs.""" + if encoder_outputs is None and input_features is not None: + encoder_outputs = self._encoder(input_features=input_features) + return super().forward(encoder_outputs=encoder_outputs, **kwargs) + + @property + def generation_config(self) -> GenerationConfig: # noqa: D102 + if not hasattr(self, "_generation_config"): + from transformers import GenerationConfig + + values: dict[str, Any] = {} + for name in ( + "decoder_start_token_id", + "bos_token_id", + "eos_token_id", + "pad_token_id", + "forced_decoder_ids", + "suppress_tokens", + "begin_suppress_tokens", + ): + value = getattr(self.config, name, None) + if value is not None: + values[name] = value + values.setdefault("max_new_tokens", self._max_dec - 1) + values.setdefault("num_beams", 1) + values.setdefault("do_sample", False) + self._generation_config = GenerationConfig(**values) + return self._generation_config + + @generation_config.setter + def generation_config(self, value: Any) -> None: + self._generation_config = value + + +__all__ = [ + "MODEL_CLASS_MAPPING", + "WhisperDecoderIOConfig", + "WhisperDecoderWrapper", + "WhisperEncoderIOConfig", + "WhisperEncoderInputGenerator", + "WhisperEncoderWrapper", + "WinMLWhisperModel", +] diff --git a/src/winml/modelkit/models/winml/encoder_decoder.py b/src/winml/modelkit/models/winml/encoder_decoder.py index 6388a42f1..1e735d5e0 100644 --- a/src/winml/modelkit/models/winml/encoder_decoder.py +++ b/src/winml/modelkit/models/winml/encoder_decoder.py @@ -161,13 +161,12 @@ def generate( # ============================================================================= -class WinMLEncoderDecoderModel(WinMLCompositeModel, GenerationMixin): - """composite model with HF GenerationMixin support. +class WinMLEncoderDecoderCore(WinMLCompositeModel): + """Shared encoder-decoder runtime independent of a generation policy. Expects sub-components ``"encoder"`` and ``"decoder"`` in - ``_SUB_MODEL_CONFIG``. Provides the full interface required by - ``GenerationMixin.generate()`` for encoder-decoder models with - static KV cache. + ``_SUB_MODEL_CONFIG``. Provides the interface required by Transformers + generation mixins for encoder-decoder models with static KV cache. Input/output names and shapes are read from ONNX I/O metadata — no model-specific names are assumed. @@ -301,7 +300,7 @@ def _validate_model_kwargs(self, model_kwargs: dict[str, Any]) -> None: } GenerationMixin._validate_model_kwargs(cast("Any", self), remaining_kwargs) - def prepare_inputs_for_generation( # type: ignore[override] # GenerationMixin's base signature differs; static-cache flow + def prepare_inputs_for_generation( self, input_ids: torch.LongTensor, past_key_values: Cache | None = None, @@ -512,3 +511,9 @@ def forward( logits=outputs["logits"], past_key_values=cache, ) + + +class WinMLEncoderDecoderModel( # type: ignore[misc] + WinMLEncoderDecoderCore, GenerationMixin +): + """Encoder-decoder runtime using Transformers' generic generation policy.""" diff --git a/src/winml/modelkit/utils/eval_utils.py b/src/winml/modelkit/utils/eval_utils.py index f85c672b4..1ec3457cf 100644 --- a/src/winml/modelkit/utils/eval_utils.py +++ b/src/winml/modelkit/utils/eval_utils.py @@ -57,6 +57,30 @@ class TaskSchema: ), ) +_AUTOMATIC_SPEECH_RECOGNITION_SCHEMA = TaskSchema( + columns=( + SchemaItem( + "input_column", + "audio bytes/path or decoded waveform with sampling rate", + default="audio", + remap_hint="", + ), + SchemaItem( + "label_column", + "reference transcript", + default="text", + remap_hint="", + ), + ), + params=( + SchemaItem("max_audio_seconds", "per-row audio duration cap", default="30"), + SchemaItem("max_new_tokens", "seq2seq generation cap", default="32"), + SchemaItem("language", "seq2seq language prompt", default="english"), + SchemaItem("generation_task", "seq2seq speech task", default="transcribe"), + ), + roles=("encoder", "decoder"), +) + _TEXT_CLASSIFICATION_SCHEMA = TaskSchema( columns=( SchemaItem( @@ -449,6 +473,7 @@ class TaskSchema: ) TASK_SCHEMAS: dict[str, TaskSchema] = { + "automatic-speech-recognition": _AUTOMATIC_SPEECH_RECOGNITION_SCHEMA, "image-classification": _IMAGE_CLASSIFICATION_SCHEMA, "text-classification": _TEXT_CLASSIFICATION_SCHEMA, "sequence-classification": _TEXT_CLASSIFICATION_SCHEMA, diff --git a/tests/unit/commands/test_build.py b/tests/unit/commands/test_build.py index 37f6bde79..79a7d3eec 100644 --- a/tests/unit/commands/test_build.py +++ b/tests/unit/commands/test_build.py @@ -18,6 +18,7 @@ import pytest from click.testing import CliRunner +from winml.modelkit.config import WinMLBuildConfig from winml.modelkit.session import EPDeviceTarget @@ -523,6 +524,79 @@ def test_no_optimize_sets_extra_kwarg(self, tmp_path: Path, mock_run_single_buil assert result.exit_code == 0, result.output assert mock_run_single_build.call_args.kwargs["extra_kwargs"].get("skip_optimize") is True + @pytest.mark.parametrize("pipeline_name", ["_build_hf_pipeline", "_build_onnx_pipeline"]) + def test_pipeline_threads_no_optimize_to_stage( + self, + pipeline_name: str, + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + """The Rich CLI pipeline must deliver ``skip_optimize`` to its stage sink.""" + import winml.modelkit.commands.build as build_module + + config = WinMLBuildConfig.from_dict( + { + "loader": {"task": "image-classification", "model_type": "resnet"}, + "export": {"opset_version": 17}, + "optim": {}, + "quant": None, + "compile": None, + } + ) + input_model = tmp_path / "input.onnx" + input_model.write_bytes(b"onnx") + observed: dict[str, bool] = {} + + def fake_optimize_stage(**kwargs): + observed["skip_optimize"] = kwargs["skip_optimize"] + kwargs["optimized_path"].write_bytes(b"onnx") + return kwargs["optimized_path"], 0.0 + + monkeypatch.setattr(build_module, "_run_optimize_stage", fake_optimize_stage) + monkeypatch.setattr( + build_module, "_run_quantize_stage", lambda **kwargs: kwargs["current_path"] + ) + monkeypatch.setattr( + build_module, "_run_compile_stage", lambda **kwargs: kwargs["current_path"] + ) + monkeypatch.setattr( + "winml.modelkit.onnx.copy_onnx_model", + lambda source, destination: Path(destination).write_bytes(Path(source).read_bytes()), + ) + + common_kwargs = { + "config": config, + "output_dir": tmp_path / "output", + "rebuild": True, + "ep": "cpu", + "device": "cpu", + "extra_kwargs": {"skip_optimize": True}, + } + if pipeline_name == "_build_hf_pipeline": + monkeypatch.setattr( + "winml.modelkit.build.hf._load_model", lambda *args, **kwargs: object() + ) + monkeypatch.setattr( + "winml.modelkit.export.export_onnx", + lambda *, output_path, **kwargs: Path(output_path).write_bytes(b"onnx"), + ) + build_module._build_hf_pipeline( + **common_kwargs, + model_id="test/model", + cache_key=None, + ) + else: + monkeypatch.setattr( + "winml.modelkit.build.common.ensure_pre_quantized_stamped", + lambda *args, **kwargs: None, + ) + build_module._build_onnx_pipeline( + **common_kwargs, + onnx_path=input_model, + ) + + assert observed["skip_optimize"] is True + def test_no_analyze_zeros_max_iterations( self, tmp_path: Path, mock_run_single_build: MagicMock ): diff --git a/tests/unit/eval/test_automatic_speech_recognition_evaluator.py b/tests/unit/eval/test_automatic_speech_recognition_evaluator.py new file mode 100644 index 000000000..977695294 --- /dev/null +++ b/tests/unit/eval/test_automatic_speech_recognition_evaluator.py @@ -0,0 +1,186 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- +"""Tests for explicit CTC and seq2seq ASR evaluation.""" + +from __future__ import annotations + +import io +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import numpy as np +import pytest +import torch + +from winml.modelkit.eval.automatic_speech_recognition_evaluator import ( + WinMLAutomaticSpeechRecognitionEvaluator, + _asr_mode, + _word_error_rate, +) +from winml.modelkit.eval.config import DatasetConfig, WinMLEvaluationConfig +from winml.modelkit.utils.eval_utils import TASK_SCHEMAS, DatasetValidationError + + +def _evaluator(*, seq2seq: bool) -> WinMLAutomaticSpeechRecognitionEvaluator: + evaluator = WinMLAutomaticSpeechRecognitionEvaluator.__new__( + WinMLAutomaticSpeechRecognitionEvaluator + ) + evaluator.config = WinMLEvaluationConfig( + model_id="test/asr", + task="automatic-speech-recognition", + dataset=DatasetConfig(path="test", samples=1), + ) + evaluator.model = MagicMock() + evaluator.model.config = SimpleNamespace(is_encoder_decoder=seq2seq) + evaluator.processor = MagicMock() + evaluator.mode = "seq2seq" if seq2seq else "ctc" + evaluator._audio_column = "audio" + evaluator._label_column = "text" + evaluator._max_audio_seconds = 30.0 + evaluator._max_new_tokens = 32 + evaluator._language = "english" + evaluator._generation_task = "transcribe" + return evaluator + + +def test_asr_schema_declares_audio_transcript_and_component_roles() -> None: + schema = TASK_SCHEMAS["automatic-speech-recognition"] + assert [item.default for item in schema.columns] == ["audio", "text"] + assert schema.roles == ("encoder", "decoder") + + +def test_mode_is_selected_from_encoder_decoder_contract() -> None: + assert _asr_mode(SimpleNamespace(config=SimpleNamespace(is_encoder_decoder=True))) == "seq2seq" + assert _asr_mode(SimpleNamespace(config=SimpleNamespace(is_encoder_decoder=False))) == "ctc" + + +def test_word_error_rate() -> None: + assert _word_error_rate("THE QUICK FOX", "the quick fox") == 0.0 + assert _word_error_rate("the fox", "the quick fox") == pytest.approx(1 / 3) + + +def test_decode_audio_mixes_resamples_and_caps() -> None: + evaluator = _evaluator(seq2seq=True) + evaluator._max_audio_seconds = 0.01 + stereo = np.ones((160, 2), dtype=np.float32) + with patch( + "scipy.signal.resample_poly", return_value=np.ones(320, dtype=np.float32) + ) as resample: + waveform = evaluator._decode_audio({"array": stereo, "sampling_rate": 8_000}) + resample.assert_called_once() + assert waveform.shape == (160,) + assert waveform.dtype == np.float32 + + +def test_decode_audio_reads_raw_bytes() -> None: + evaluator = _evaluator(seq2seq=True) + decoded = np.ones(160, dtype=np.float32) + with patch("soundfile.read", return_value=(decoded, 16_000)) as read_audio: + waveform = evaluator._decode_audio({"bytes": b"flac", "path": None}) + + source = read_audio.call_args.args[0] + assert isinstance(source, io.BytesIO) + assert source.read() == b"flac" + assert np.array_equal(waveform, decoded) + + +def test_ctc_compute_decodes_frame_logits_and_accounts_exactly() -> None: + evaluator = _evaluator(seq2seq=False) + evaluator.data = [{"audio": {"array": np.ones(160), "sampling_rate": 16_000}, "text": "hi"}] + evaluator.processor.return_value = {"input_values": torch.ones(1, 160)} + evaluator.processor.batch_decode.return_value = ["hi"] + evaluator.model.return_value = SimpleNamespace(logits=torch.tensor([[[0.0, 1.0], [1.0, 0.0]]])) + + result = evaluator.compute() + + evaluator.model.assert_called_once() + evaluator.model.generate.assert_not_called() + assert result == { + "wer": 0.0, + "asr_mode": "ctc", + "requested_samples": 1, + "processed_samples": 1, + "skipped_samples": 0, + "max_audio_seconds": 30.0, + "max_new_tokens": 0, + "num_beams": 0, + "return_sequences": 1, + } + + +def test_compute_uses_corpus_word_error_rate() -> None: + evaluator = _evaluator(seq2seq=False) + evaluator.config.dataset.samples = 2 + evaluator.data = [ + {"audio": {"array": np.ones(160), "sampling_rate": 16_000}, "text": "one"}, + { + "audio": {"array": np.ones(160), "sampling_rate": 16_000}, + "text": "two three four", + }, + ] + evaluator.processor.return_value = {"input_values": torch.ones(1, 160)} + evaluator.processor.batch_decode.side_effect = [["wrong"], ["two three four"]] + evaluator.model.return_value = SimpleNamespace(logits=torch.ones(1, 1, 1)) + + result = evaluator.compute() + + assert result["wer"] == 0.25 + assert result["processed_samples"] == 2 + + +def test_compute_scores_empty_hypothesis_as_deletions() -> None: + evaluator = _evaluator(seq2seq=False) + evaluator.config.dataset.samples = 2 + evaluator.data = [ + { + "audio": {"array": np.ones(160), "sampling_rate": 16_000}, + "text": "hello world", + }, + { + "audio": {"array": np.ones(160), "sampling_rate": 16_000}, + "text": "recognized", + }, + ] + evaluator.processor.return_value = {"input_values": torch.ones(1, 160)} + evaluator.processor.batch_decode.side_effect = [[""], ["recognized"]] + evaluator.model.return_value = SimpleNamespace(logits=torch.ones(1, 1, 1)) + + result = evaluator.compute() + + assert result["wer"] == pytest.approx(2 / 3) + assert result["processed_samples"] == 2 + + +def test_seq2seq_compute_uses_bounded_generation() -> None: + evaluator = _evaluator(seq2seq=True) + evaluator.data = [{"audio": {"array": np.ones(160), "sampling_rate": 16_000}, "text": "hello"}] + evaluator.processor.return_value = {"input_features": torch.ones(1, 80, 3000)} + evaluator.processor.get_decoder_prompt_ids = MagicMock() + evaluator.processor.batch_decode.return_value = ["hello"] + evaluator.model.generation_config = SimpleNamespace() + evaluator.model.generate.return_value = torch.tensor([[1, 2]]) + + result = evaluator.compute() + + generation_config = evaluator.model.generate.call_args.kwargs["generation_config"] + assert generation_config.max_new_tokens == 32 + assert generation_config.num_beams == 1 + assert generation_config.do_sample is False + assert generation_config.language == "english" + assert generation_config.task == "transcribe" + assert generation_config.return_timestamps is False + assert result["asr_mode"] == "seq2seq" + assert result["processed_samples"] == 1 + + +def test_compute_fails_closed_on_missing_or_empty_data() -> None: + evaluator = _evaluator(seq2seq=False) + evaluator.data = [{"audio": None, "text": "reference"}] + with pytest.raises(DatasetValidationError, match="row 0"): + evaluator.compute() + + evaluator.data = [] + with pytest.raises(DatasetValidationError, match="no usable rows"): + evaluator.compute() diff --git a/tests/unit/models/whisper/test_onnx_config.py b/tests/unit/models/whisper/test_onnx_config.py new file mode 100644 index 000000000..b35dbf65c --- /dev/null +++ b/tests/unit/models/whisper/test_onnx_config.py @@ -0,0 +1,152 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- +"""Tests for split Whisper encoder/decoder support.""" + +from __future__ import annotations + +from types import SimpleNamespace + +from optimum.exporters.tasks import TasksManager +from transformers import GenerationConfig, WhisperConfig +from transformers.models.whisper.generation_whisper import WhisperGenerationMixin + +from winml.modelkit.loader.resolution import ( + _composite_components_for_task, + resolve_composite, +) +from winml.modelkit.models.hf.whisper import ( + WhisperDecoderIOConfig, + WhisperEncoderInputGenerator, + WhisperEncoderIOConfig, + WinMLWhisperModel, + _WhisperDecoderNormalizedConfig, +) +from winml.modelkit.models.winml.composite_model import COMPOSITE_MODEL_REGISTRY +from winml.modelkit.models.winml.encoder_decoder import EncoderDecoderInputGenerator +from winml.modelkit.models.winml.kv_cache import PastKeyValueInputGenerator, WinMLStaticCache + + +def _config() -> WhisperConfig: + return WhisperConfig( + vocab_size=128, + num_mel_bins=80, + d_model=32, + encoder_layers=2, + decoder_layers=3, + encoder_attention_heads=2, + decoder_attention_heads=4, + encoder_ffn_dim=64, + decoder_ffn_dim=64, + max_source_positions=20, + max_target_positions=16, + pad_token_id=0, + bos_token_id=1, + eos_token_id=2, + decoder_start_token_id=1, + ) + + +def test_encoder_registration_and_contract() -> None: + constructor = TasksManager.get_exporter_config_constructor( + model_type="whisper", + exporter="onnx", + task="feature-extraction", + library_name="transformers", + ) + assert constructor.func is WhisperEncoderIOConfig + onnx_config = WhisperEncoderIOConfig(_config(), task="feature-extraction") + assert onnx_config.inputs == {"input_features": {0: "batch_size"}} + assert onnx_config.outputs == {"encoder_hidden_states": {0: "batch_size"}} + + +def test_encoder_dummy_input_uses_fixed_audio_geometry() -> None: + onnx_config = WhisperEncoderIOConfig(_config(), task="feature-extraction") + inputs = onnx_config.generate_dummy_inputs(framework="pt") + assert inputs["input_features"].shape == (1, 80, 40) + assert (WhisperEncoderInputGenerator,) == WhisperEncoderIOConfig.DUMMY_INPUT_GENERATOR_CLASSES + + +def test_decoder_contract_uses_decoder_dimensions() -> None: + config = _config() + normalized = _WhisperDecoderNormalizedConfig(config) + assert normalized.num_layers == 3 + assert normalized.num_attention_heads == 4 + assert normalized.head_dim == 8 + assert normalized.max_cache_len == 16 + assert normalized.sequence_length == 20 + + onnx_config = WhisperDecoderIOConfig(config, task="text2text-generation") + assert ( + EncoderDecoderInputGenerator, + PastKeyValueInputGenerator, + ) == WhisperDecoderIOConfig.DUMMY_INPUT_GENERATOR_CLASSES + assert set(onnx_config.inputs).issuperset( + { + "decoder_input_ids", + "encoder_hidden_states", + "decoder_attention_mask", + "cache_position", + "past_0_key", + "past_2_value", + } + ) + assert "past_3_key" not in onnx_config.inputs + assert set(onnx_config.outputs).issuperset({"logits", "present_0_key", "present_2_value"}) + + +def test_asr_composite_registration_and_detection_bridge() -> None: + expected = { + "encoder": "feature-extraction", + "decoder": "text2text-generation", + } + assert ( + COMPOSITE_MODEL_REGISTRY[("whisper", "automatic-speech-recognition")] is WinMLWhisperModel + ) + assert resolve_composite("whisper", "automatic-speech-recognition") == expected + assert _composite_components_for_task("whisper", "automatic-speech-recognition") == expected + assert WinMLWhisperModel.main_input_name == "input_features" + assert WinMLWhisperModel.get_cache_class() is WinMLStaticCache + assert issubclass(WinMLWhisperModel, WhisperGenerationMixin) + assert WinMLWhisperModel.generate is WhisperGenerationMixin.generate + + +def test_generation_config_preserves_whisper_tokens() -> None: + model = object.__new__(WinMLWhisperModel) + model.config = _config() + model._max_dec = 16 + generation_config = model.generation_config + assert generation_config.decoder_start_token_id == 1 + assert generation_config.eos_token_id == 2 + assert generation_config.max_new_tokens == 15 + + +def test_runtime_loads_generation_metadata_and_derives_input_stride(monkeypatch) -> None: + config = _config() + config._name_or_path = "test/whisper" + loaded_generation_config = GenerationConfig(decoder_start_token_id=7) + monkeypatch.setattr( + GenerationConfig, + "from_pretrained", + lambda model_name_or_path: loaded_generation_config, + ) + encoder = SimpleNamespace( + io_config={ + "input_names": ["input_features"], + "input_shapes": [[1, 80, 40]], + } + ) + decoder = SimpleNamespace( + io_config={ + "input_names": ["decoder_input_ids", "past_0_key"], + "input_shapes": [[1, 1], [1, 4, 16, 8]], + "input_types": ["int64", "float32"], + } + ) + + model = WinMLWhisperModel({"encoder": encoder, "decoder": decoder}, config) + + assert model.generation_config is loaded_generation_config + assert model.model.encoder.conv1.stride == (1,) + assert model.model.encoder.conv2.stride == (2,)