Add provider selection for gpt2/convert_to_onnx.py (#13982)

Allows the user to select from supported backends for gpt2/convert_to_onnx.py. Default behavior is preserved if no provider is selected. This allows the ROCm EP to be selected.
This commit is contained in:
Joseph Groenenboom 2022-12-21 21:41:09 -06:00 committed by GitHub
parent a170e40fbb
commit baba312e30
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23

View file

@ -85,6 +85,14 @@ def parse_arguments(argv=None):
parser.add_argument("--use_gpu", required=False, action="store_true", help="use GPU for inference")
parser.set_defaults(use_gpu=False)
parser.add_argument(
"--provider",
required=False,
default=None,
choices=["dml", "rocm", "migraphx", "cuda", "tensorrt"],
help="use dml, rocm, cuda, tensorrt or migraphx for respective backend",
)
parser.add_argument(
"--tolerance",
required=False,
@ -420,7 +428,9 @@ def main(argv=None, experiment_name="", run_id=0, csv_filename="gpt2_parity_resu
logger.info(f"Output path: {output_path}")
model_size_in_MB = int(get_onnx_model_size(output_path, args.use_external_data_format) / 1024 / 1024)
session = create_onnxruntime_session(output_path, args.use_gpu, enable_all_optimization=True, verbose=args.verbose)
session = create_onnxruntime_session(
output_path, args.use_gpu, args.provider, enable_all_optimization=True, verbose=args.verbose
)
if args.model_class == "GPT2LMHeadModel" and session is not None:
parity_result = gpt2helper.test_parity(
session,