From ace41b80647a627805c766ef6b05e78a1e5cd206 Mon Sep 17 00:00:00 2001 From: Tianlei Wu Date: Thu, 23 Jul 2020 11:35:41 -0700 Subject: [PATCH] Force return_tuple=True to handle transformers breaking change of output format. (#4599) --- onnxruntime/python/tools/transformers/benchmark_gpt2.py | 2 ++ onnxruntime/python/tools/transformers/convert_to_onnx.py | 5 ++++- onnxruntime/python/tools/transformers/onnx_exporter.py | 2 ++ 3 files changed, 8 insertions(+), 1 deletion(-) diff --git a/onnxruntime/python/tools/transformers/benchmark_gpt2.py b/onnxruntime/python/tools/transformers/benchmark_gpt2.py index f5fa75550f..1e281822f8 100644 --- a/onnxruntime/python/tools/transformers/benchmark_gpt2.py +++ b/onnxruntime/python/tools/transformers/benchmark_gpt2.py @@ -126,6 +126,8 @@ def main(): model_class = MODEL_CLASSES[args.model_class][0] config = AutoConfig.from_pretrained(args.model_name, torchscript=args.torchscript, cache_dir=cache_dir) + if hasattr(config, 'return_tuple'): + config.return_tuple = True model = model_class.from_pretrained(args.model_name, config=config, cache_dir=cache_dir) # This scirpt does not support float16 for PyTorch. diff --git a/onnxruntime/python/tools/transformers/convert_to_onnx.py b/onnxruntime/python/tools/transformers/convert_to_onnx.py index 47a8ac401b..809457bade 100644 --- a/onnxruntime/python/tools/transformers/convert_to_onnx.py +++ b/onnxruntime/python/tools/transformers/convert_to_onnx.py @@ -111,7 +111,10 @@ def main(): assert not args.use_gpu, "quantization only supports CPU" model_class = MODEL_CLASSES[args.model_class][0] - model = model_class.from_pretrained(args.model_name_or_path, cache_dir=cache_dir) + config = AutoConfig.from_pretrained(args.model_name_or_path, cache_dir=cache_dir) + if hasattr(config, 'return_tuple'): + config.return_tuple = True + model = model_class.from_pretrained(args.model_name_or_path, config=config, cache_dir=cache_dir) device = torch.device("cuda:0" if args.use_gpu else "cpu") model.eval().to(device) diff --git a/onnxruntime/python/tools/transformers/onnx_exporter.py b/onnxruntime/python/tools/transformers/onnx_exporter.py index 0622223ecc..61e2f93124 100644 --- a/onnxruntime/python/tools/transformers/onnx_exporter.py +++ b/onnxruntime/python/tools/transformers/onnx_exporter.py @@ -167,6 +167,8 @@ def export_onnx_model(model_name, opset_version, use_external_data_format, model use_gpu, precision, optimize_onnx, validate_onnx, use_raw_attention_mask, overwrite, model_fusion_statistics): config = AutoConfig.from_pretrained(model_name, cache_dir=cache_dir) + if hasattr(config, 'return_tuple'): + config.return_tuple = True model = load_pretrained_model(model_name, config=config, cache_dir=cache_dir) model.cpu()