mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
add quantization support for whisper (#15589)
### Description <!-- Describe your changes. --> Add dynamic quantization support for whisper model. There are 3 options to try out: - quantize_embedding_layer: enable to quantize embedding layer of decoder model or not - quantize_per_channel: enable to quantize per channel for Gemm or MatMul - quantize_reduce_range: use 7bit to quantize MatMul or Gemm. Use when hitting accuracy issue on x64 cpus without VNNI.
This commit is contained in:
parent
4b74cb1741
commit
373f912e51
1 changed files with 58 additions and 16 deletions
|
|
@ -14,6 +14,8 @@ import torch
|
|||
from whisper_chain import chain_model
|
||||
from whisper_helper import PRETRAINED_WHISPER_MODELS, WhisperHelper
|
||||
|
||||
from onnxruntime import quantization
|
||||
|
||||
sys.path.append(os.path.join(os.path.dirname(__file__), "..", ".."))
|
||||
from benchmark_helper import Precision, create_onnxruntime_session, prepare_environment, setup_logger # noqa: E402
|
||||
|
||||
|
|
@ -67,8 +69,8 @@ def parse_arguments():
|
|||
required=False,
|
||||
type=Precision,
|
||||
default=Precision.FLOAT32,
|
||||
choices=[Precision.FLOAT32, Precision.FLOAT16],
|
||||
help="Precision of model to run. fp32 for full precision, fp16 for half precision",
|
||||
choices=[Precision.FLOAT32, Precision.FLOAT16, Precision.INT8],
|
||||
help="Precision of model to run. fp32 for full precision, fp16 for half precision, int8 for quantization",
|
||||
)
|
||||
|
||||
parser.add_argument("--verbose", required=False, action="store_true")
|
||||
|
|
@ -134,6 +136,27 @@ def parse_arguments():
|
|||
help="default name is whisper_beamsearch.onnx.",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--quantize_embedding_layer",
|
||||
required=False,
|
||||
action="store_true",
|
||||
help="Produce beam search model with chained encdecinit and decoder.",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--quantize_per_channel",
|
||||
required=False,
|
||||
action="store_true",
|
||||
help="Produce beam search model with chained encdecinit and decoder.",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--quantize_reduce_range",
|
||||
required=False,
|
||||
action="store_true",
|
||||
help="Produce beam search model with chained encdecinit and decoder.",
|
||||
)
|
||||
|
||||
parser.add_argument("--no_repeat_ngram_size", type=int, default=3, help="default to 3")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
|
@ -155,6 +178,9 @@ def export_onnx_models(
|
|||
overwrite: bool = False,
|
||||
disable_auto_mixed_precision: bool = False,
|
||||
use_int32_inputs: bool = True,
|
||||
quantize_embedding_layer: bool = False,
|
||||
quantize_per_channel: bool = False,
|
||||
quantize_reduce_range: bool = False,
|
||||
):
|
||||
device = torch.device("cuda:0" if use_gpu else "cpu")
|
||||
|
||||
|
|
@ -202,17 +228,33 @@ def export_onnx_models(
|
|||
)
|
||||
|
||||
if overwrite or not os.path.exists(output_path):
|
||||
logger.info(f"Optimizing model to {output_path}")
|
||||
WhisperHelper.optimize_onnx(
|
||||
onnx_path,
|
||||
output_path,
|
||||
precision == Precision.FLOAT16,
|
||||
config.encoder_attention_heads,
|
||||
config.d_model,
|
||||
use_external_data_format,
|
||||
auto_mixed_precision=not disable_auto_mixed_precision,
|
||||
use_gpu=use_gpu,
|
||||
)
|
||||
if optimize_onnx:
|
||||
logger.info(f"Optimizing model to {output_path}")
|
||||
WhisperHelper.optimize_onnx(
|
||||
onnx_path,
|
||||
output_path,
|
||||
precision == Precision.FLOAT16,
|
||||
config.encoder_attention_heads,
|
||||
config.d_model,
|
||||
use_external_data_format,
|
||||
auto_mixed_precision=not disable_auto_mixed_precision,
|
||||
use_gpu=use_gpu,
|
||||
)
|
||||
onnx_path = output_path
|
||||
|
||||
if precision == Precision.INT8:
|
||||
quantization.quantize_dynamic(
|
||||
onnx_path,
|
||||
output_path,
|
||||
op_types_to_quantize=["MatMul", "Gemm", "Gather"]
|
||||
if quantize_embedding_layer
|
||||
else ["MatMul", "Gemm"],
|
||||
use_external_data_format=use_external_data_format,
|
||||
per_channel=quantize_per_channel,
|
||||
reduce_range=quantize_reduce_range,
|
||||
optimize_model=False,
|
||||
extra_options={"MatMulConstBOnly": True},
|
||||
)
|
||||
else:
|
||||
logger.info(f"Skip optimizing: existed ONNX model {onnx_path}")
|
||||
else:
|
||||
|
|
@ -246,9 +288,6 @@ def main():
|
|||
output_dir = args.output if not args.output.endswith(".onnx") else os.path.dirname(args.output)
|
||||
prepare_environment(cache_dir, output_dir, args.use_gpu)
|
||||
|
||||
if args.precision != Precision.FLOAT32:
|
||||
assert args.optimize_onnx, "fp16/int8 requires --optimize_onnx"
|
||||
|
||||
if args.precision == Precision.FLOAT16:
|
||||
assert args.use_gpu, "fp16 requires --use_gpu"
|
||||
|
||||
|
|
@ -269,6 +308,9 @@ def main():
|
|||
args.overwrite,
|
||||
args.disable_auto_mixed_precision,
|
||||
not args.use_int64_inputs,
|
||||
args.quantize_embedding_layer,
|
||||
args.quantize_per_channel,
|
||||
args.quantize_reduce_range,
|
||||
)
|
||||
|
||||
if args.chain_model:
|
||||
|
|
|
|||
Loading…
Reference in a new issue