From 994c2ed420f4764352404f5ccd91738f61bc1182 Mon Sep 17 00:00:00 2001 From: Xiaoyu Liu Date: Thu, 29 Apr 2021 13:19:56 -0700 Subject: [PATCH] GPT2 one step beam search update with configuration support (#7425) * check in early stop search as separate type * rename to beam search configurations * update do sample configuration flag help * rename to configurable search step * add option groups * add more unit tests Co-authored-by: Xiaoyu Liu --- .../tools/transformers/benchmark_gpt2.py | 144 ++++- .../tools/transformers/convert_to_onnx.py | 85 ++- .../transformers/gpt2_beamsearch_helper.py | 571 +++++++++++------- .../transformers/gpt2_beamsearch_tester.py | 234 ++++--- .../tools/transformers/test/test_gpt2.py | 31 + 5 files changed, 679 insertions(+), 386 deletions(-) diff --git a/onnxruntime/python/tools/transformers/benchmark_gpt2.py b/onnxruntime/python/tools/transformers/benchmark_gpt2.py index ed5d3de6c8..a0407b30ba 100644 --- a/onnxruntime/python/tools/transformers/benchmark_gpt2.py +++ b/onnxruntime/python/tools/transformers/benchmark_gpt2.py @@ -18,7 +18,8 @@ import torch import onnx from packaging import version from transformers import AutoConfig -from gpt2_helper import Gpt2Helper, MODEL_CLASSES, DEFAULT_TOLERANCE, PRETRAINED_GPT2_MODELS +from gpt2_helper import DEFAULT_TOLERANCE, PRETRAINED_GPT2_MODELS +from gpt2_beamsearch_helper import Gpt2HelperFactory, MODEL_CLASSES from quantize_helper import QuantizeHelper from benchmark_helper import create_onnxruntime_session, setup_logger, prepare_environment, Precision @@ -84,6 +85,7 @@ def parse_arguments(argv=None): parser.set_defaults(torchscript=False) parser.add_argument('-b', '--batch_sizes', nargs='+', type=int, default=[1], help="batch size") + parser.add_argument('--beam_size', type=int, default=4, help='Beam size if greedy/top-p/top-k sampling is needed') parser.add_argument('--sequence_lengths', nargs='+', @@ -108,6 +110,40 @@ def parse_arguments(argv=None): parser.add_argument('--verbose', required=False, action='store_true') parser.set_defaults(verbose=False) + search_option_group = parser.add_argument_group("configurable one step search options") + + search_option_group.add_argument('--ignore_eos', + type=bool, + default=False, + help='If ignore end of sentence token in model inference.') + search_option_group.add_argument('--repetition_penalty', + type=float, + default=1, + help='Positive. >1 to penalize and <1 to encorage.') + search_option_group.add_argument('--temperature', + type=float, + default=1, + help='Softmax temperature for output logits.') + search_option_group.add_argument('--excluded_token_ids', + required=False, + nargs='+', + type=float, + help='A list of token ids to be excluded in inference.') + search_option_group.add_argument('--length_penalty', + type=float, + default=1, + help='Positive. >1 to penalize and <1 to encorage short sentence.') + + sampling_option_group = parser.add_argument_group("one step sampling options") + sampling_option_group.add_argument('--do_sample', + action='store_true', + help='If to do sampling instead of beam search or greedy.') + sampling_option_group.add_argument('--do_sample_top_p', + type=float, + default=0.95, + help='Nuclear/top-p sampling accumulation probability.') + sampling_option_group.add_argument('--do_sample_top_k', type=int, default=0, help='Use top-k if non-zero.') + args = parser.parse_args(argv) return args @@ -115,7 +151,8 @@ def parse_arguments(argv=None): def main(args): from transformers import __version__ as transformers_version - if version.parse(transformers_version) < version.parse("3.1.0"): # past_key_values name does not exist in 3.0.2 or older + if version.parse(transformers_version) < version.parse( + "3.1.0"): # past_key_values name does not exist in 3.0.2 or older raise RuntimeError("This tool requires transformers 3.1.0 or later.") logger.info(f"Arguments:{args}") @@ -133,9 +170,37 @@ def main(args): prepare_environment(cache_dir, output_dir, args.use_gpu) model_class = MODEL_CLASSES[args.model_class][0] + if args.model_class == "GPT2LMHeadModel_BeamSearchStep": + model_type = "beam_search_step" + elif args.model_class == "GPT2LMHeadModel_ConfigurableOneStepSearch": + model_type = "configurable_one_step_search" + else: + model_type = "default" + gpt2helper = Gpt2HelperFactory.create_helper(model_type) config = AutoConfig.from_pretrained(args.model_name_or_path, torchscript=args.torchscript, cache_dir=cache_dir) - model = model_class.from_pretrained(args.model_name_or_path, config=config, cache_dir=cache_dir) + if model_type == 'beam_search_step': + model = model_class.from_pretrained(args.model_name_or_path, + config=config, + batch_size=1, + beam_size=args.beam_size, + cache_dir=cache_dir) + elif model_type == 'configurable_one_step_search': + model = model_class.from_pretrained(args.model_name_or_path, + config=config, + batch_size=1, + beam_size=args.beam_size, + ignore_eos=args.ignore_eos, + temperature=args.temperature, + repetition_penalty=args.repetition_penalty, + excluded_token_ids=args.excluded_token_ids, + length_penalty=args.length_penalty, + do_sample=args.do_sample, + do_sample_top_p=args.do_sample_top_p, + do_sample_top_k=args.do_sample_top_k, + cache_dir=cache_dir) + else: + model = model_class.from_pretrained(args.model_name_or_path, config=config, cache_dir=cache_dir) # This scirpt does not support float16 for PyTorch. #if args.float16: @@ -144,7 +209,7 @@ def main(args): device = torch.device("cuda:0" if args.use_gpu else "cpu") model.to(device) use_external_data_format = (config.n_layer > 24) #TODO: find a way to check model size > 2GB - onnx_model_paths = Gpt2Helper.get_onnx_paths(output_dir, + onnx_model_paths = gpt2helper.get_onnx_paths(output_dir, args.model_name_or_path, args.model_class, has_past=True, @@ -152,7 +217,7 @@ def main(args): onnx_model_path = onnx_model_paths["raw"] use_padding = MODEL_CLASSES[args.model_class][2] - Gpt2Helper.export_onnx(model, + gpt2helper.export_onnx(model, device, onnx_model_path, args.verbose, @@ -162,7 +227,7 @@ def main(args): if args.optimize_onnx or args.precision != Precision.FLOAT32: onnx_model_path = onnx_model_paths[str(args.precision) if args.precision != Precision.INT8 else 'fp32'] - Gpt2Helper.optimize_onnx(onnx_model_paths["raw"], onnx_model_path, args.precision == Precision.FLOAT16, + gpt2helper.optimize_onnx(onnx_model_paths["raw"], onnx_model_path, args.precision == Precision.FLOAT16, model.config.num_attention_heads, model.config.hidden_size, use_external_data_format) if args.precision == Precision.INT8: @@ -173,7 +238,7 @@ def main(args): onnx_model_path = onnx_model_paths["int8"] if args.torchscript: - model = Gpt2Helper.torchscript(model, + model = gpt2helper.torchscript(model, config, device, has_position_ids=use_padding, @@ -188,9 +253,16 @@ def main(args): return # Allocate output buffers for IO Binding - max_output_shapes = Gpt2Helper.get_output_shapes(max(args.batch_sizes), max(args.past_sequence_lengths), - max(args.sequence_lengths), config, args.model_class) - output_buffers = Gpt2Helper.get_output_buffers(max_output_shapes, device, args.precision == Precision.FLOAT16) + if model_type == 'beam_search_step' or model_type == 'configurable_one_step_search': + max_output_shapes = gpt2helper.get_output_shapes(max(args.batch_sizes), max(args.past_sequence_lengths), + max(args.past_sequence_lengths), max(args.sequence_lengths), 4, + 0, config, args.model_class) + output_buffers = gpt2helper.get_output_buffers(max_output_shapes, device, args.precision == Precision.FLOAT16) + + else: + max_output_shapes = gpt2helper.get_output_shapes(max(args.batch_sizes), max(args.past_sequence_lengths), + max(args.sequence_lengths), config, args.model_class) + output_buffers = gpt2helper.get_output_buffers(max_output_shapes, device, args.precision == Precision.FLOAT16) csv_filename = args.result_csv or "benchmark_result_{}.csv".format(datetime.now().strftime("%Y%m%d-%H%M%S")) with open(csv_filename, mode="a", newline='') as csv_file: @@ -209,25 +281,41 @@ def main(args): logger.debug( f"Running test for batch_size={batch_size} sequence_length={sequence_length} past_sequence_length={past_sequence_length}..." ) - dummy_inputs = Gpt2Helper.get_dummy_inputs(batch_size, - past_sequence_length, - sequence_length, - config.num_attention_heads, - config.hidden_size, - config.n_layer, - config.vocab_size, - device, - float16=(args.precision == Precision.FLOAT16), - has_position_ids=use_padding, - has_attention_mask=use_padding) - output_shapes = Gpt2Helper.get_output_shapes(batch_size, past_sequence_length, sequence_length, - config, args.model_class) + if model_type == 'beam_search_step' or model_type == 'configurable_one_step_search': + dummy_inputs = gpt2helper.get_dummy_inputs(batch_size, + past_sequence_length, + sequence_length, + config.num_attention_heads, + config.hidden_size, + config.n_layer, + config.vocab_size, + device, + float16=(args.precision == Precision.FLOAT16), + has_position_ids=use_padding, + has_attention_mask=use_padding) + output_shapes = gpt2helper.get_output_shapes(batch_size, past_sequence_length, + past_sequence_length, sequence_length, 4, 0, + config, args.model_class) + else: + dummy_inputs = gpt2helper.get_dummy_inputs(batch_size, + past_sequence_length, + sequence_length, + config.num_attention_heads, + config.hidden_size, + config.n_layer, + config.vocab_size, + device, + float16=(args.precision == Precision.FLOAT16), + has_position_ids=use_padding, + has_attention_mask=use_padding) + output_shapes = gpt2helper.get_output_shapes(batch_size, past_sequence_length, sequence_length, + config, args.model_class) try: outputs, torch_latency = Gpt2Helper.pytorch_inference(model, dummy_inputs, args.test_times) ort_outputs, ort_latency = Gpt2Helper.onnxruntime_inference(session, dummy_inputs, args.test_times) - ort_io_outputs, ort_io_latency = Gpt2Helper.onnxruntime_inference_with_binded_io( + ort_io_outputs, ort_io_latency = gpt2helper.onnxruntime_inference_with_binded_io( session, dummy_inputs, output_buffers, @@ -237,8 +325,9 @@ def main(args): include_copy_output_latency=args.include_copy_output_latency) if args.validate_onnx: - if Gpt2Helper.compare_outputs(outputs, + if gpt2helper.compare_outputs(outputs, ort_outputs, + model_class, rtol=DEFAULT_TOLERANCE[args.precision], atol=DEFAULT_TOLERANCE[args.precision]): logger.info( @@ -250,8 +339,9 @@ def main(args): for output in ort_io_outputs: copy_outputs.append(output.cpu().numpy()) - if Gpt2Helper.compare_outputs(outputs, + if gpt2helper.compare_outputs(outputs, copy_outputs, + model_class, rtol=DEFAULT_TOLERANCE[args.precision], atol=DEFAULT_TOLERANCE[args.precision]): logger.info( @@ -284,7 +374,7 @@ def main(args): return csv_filename -if __name__ == '__main__': +if __name__ == '__main__': args = parse_arguments() setup_logger(args.verbose) main(args) diff --git a/onnxruntime/python/tools/transformers/convert_to_onnx.py b/onnxruntime/python/tools/transformers/convert_to_onnx.py index 013ebe66bf..89be710762 100644 --- a/onnxruntime/python/tools/transformers/convert_to_onnx.py +++ b/onnxruntime/python/tools/transformers/convert_to_onnx.py @@ -99,9 +99,41 @@ def parse_arguments(): parser.add_argument('-e', '--use_external_data_format', required=False, action='store_true') parser.set_defaults(use_external_data_format=False) + parser.add_argument('--beam_size', type=int, default=4, help='Beam size if greedy/top-p/top-k sampling is needed') - parser.add_argument('--batch_size', required=False, type=int, default=1, help='Batch size for GPT model with beam search') - parser.add_argument('--beam_size', required=False, type=int, default=4, help='Beam size for beam search') + search_option_group = parser.add_argument_group("configurable one step search options") + + search_option_group.add_argument('--ignore_eos', + type=bool, + default=False, + help='If ignore end of sentence token in model inference.') + search_option_group.add_argument('--repetition_penalty', + type=float, + default=1, + help='Positive. >1 to penalize and <1 to encorage.') + search_option_group.add_argument('--temperature', + type=float, + default=1, + help='Softmax temperature for output logits.') + search_option_group.add_argument('--excluded_token_ids', + required=False, + nargs='+', + type=float, + help='A list of token ids to be excluded in inference.') + search_option_group.add_argument('--length_penalty', + type=float, + default=1, + help='Positive. >1 to penalize and <1 to encorage short sentence.') + + sampling_option_group = parser.add_argument_group("one step sampling options") + sampling_option_group.add_argument('--do_sample', + action='store_true', + help='If to do sampling instead of beam search or greedy.') + sampling_option_group.add_argument('--do_sample_top_p', + type=float, + default=0.95, + help='Nuclear/top-p sampling accumulation probability.') + sampling_option_group.add_argument('--do_sample_top_k', type=int, default=0, help='Use top-k if non-zero.') args = parser.parse_args() @@ -110,7 +142,8 @@ def parse_arguments(): def main(): from transformers import __version__ as transformers_version - if version.parse(transformers_version) < version.parse("3.1.0"): # past_key_values name does not exist in 3.0.2 or older + if version.parse(transformers_version) < version.parse( + "3.1.0"): # past_key_values name does not exist in 3.0.2 or older raise RuntimeError("This tool requires transformers 3.1.0 or later.") args = parse_arguments() @@ -138,12 +171,36 @@ def main(): assert not args.output.endswith('.onnx'), "output shall be a directory for --use_external_data_format" model_class = MODEL_CLASSES[args.model_class][0] - model_type = "beam_search_step" if args.model_class == "GPT2LMHeadModel_BeamSearchStep" else "default" + if args.model_class == "GPT2LMHeadModel_BeamSearchStep": + model_type = "beam_search_step" + elif args.model_class == "GPT2LMHeadModel_ConfigurableOneStepSearch": + model_type = "configurable_one_step_search" + else: + model_type = "default" + gpt2helper = Gpt2HelperFactory.create_helper(model_type) gpt2tester = Gpt2TesterFactory.create_tester(model_type) config = AutoConfig.from_pretrained(args.model_name_or_path, cache_dir=cache_dir) if model_type == 'beam_search_step': - model = model_class.from_pretrained(args.model_name_or_path, config=config, batch_size=args.batch_size, beam_size=args.beam_size, cache_dir=cache_dir) + model = model_class.from_pretrained(args.model_name_or_path, + config=config, + batch_size=1, + beam_size=args.beam_size, + cache_dir=cache_dir) + elif model_type == 'configurable_one_step_search': + model = model_class.from_pretrained(args.model_name_or_path, + config=config, + batch_size=1, + beam_size=args.beam_size, + ignore_eos=args.ignore_eos, + temperature=args.temperature, + repetition_penalty=args.repetition_penalty, + excluded_token_ids=args.excluded_token_ids, + length_penalty=args.length_penalty, + do_sample=args.do_sample, + do_sample_top_p=args.do_sample_top_p, + do_sample_top_k=args.do_sample_top_k, + cache_dir=cache_dir) else: model = model_class.from_pretrained(args.model_name_or_path, config=config, cache_dir=cache_dir) @@ -239,20 +296,16 @@ def main(): else: inputs = {"input_ids": input_ids} - if model_type == "beam_search_step": + if model_type == "beam_search_step" or model_type == "configurable_one_step_search": beam_select_idx = torch.zeros([1, input_ids.shape[0]]).long() input_log_probs = torch.zeros([input_ids.shape[0], 1]) - input_unfinished_sents = torch.ones( - [input_ids.shape[0], 1], dtype=torch.bool - ) - inputs.update( - { - "beam_select_idx": beam_select_idx, - "input_log_probs": input_log_probs, - "input_unfinished_sents": input_unfinished_sents, - } - ) + input_unfinished_sents = torch.ones([input_ids.shape[0], 1], dtype=torch.bool) + inputs.update({ + "beam_select_idx": beam_select_idx, + "input_log_probs": input_log_probs, + "input_unfinished_sents": input_unfinished_sents, + }) test_inputs.append(inputs) diff --git a/onnxruntime/python/tools/transformers/gpt2_beamsearch_helper.py b/onnxruntime/python/tools/transformers/gpt2_beamsearch_helper.py index 3757397f98..a3fa8ec875 100644 --- a/onnxruntime/python/tools/transformers/gpt2_beamsearch_helper.py +++ b/onnxruntime/python/tools/transformers/gpt2_beamsearch_helper.py @@ -20,20 +20,24 @@ from gpt2_helper import Gpt2Helper, Gpt2Inputs, GPT2ModelNoPastState, MyGPT2Mode logger = logging.getLogger(__name__) +BIG_NEG = -1e4 + + class Gpt2HelperFactory: @staticmethod def create_helper(helper_type="default"): helpers = { "default": Gpt2Helper, "beam_search_step": Gpt2BeamSearchHelper, + "configurable_one_step_search": Gpt2BeamSearchHelper, } w = helpers[helper_type] return w + class GPT2LMHeadModel_BeamSearchStep(GPT2LMHeadModel): """Here we wrap a class for Onnx model conversion for GPT2LMHeadModel with past state and one step beam search.""" - def __init__(self, config, batch_size, beam_size): super().__init__(config) self.config.batch_size = batch_size @@ -61,64 +65,48 @@ class GPT2LMHeadModel_BeamSearchStep(GPT2LMHeadModel): return_dict=False, ) logits_flat, present_flat = MyGPT2Model.post_process(result, self.config.n_layer) - next_token_logits = logits_flat[:, -1].view( - self.config.batch_size, -1, logits_flat.size(-1) - ) + next_token_logits = logits_flat[:, -1].view(self.config.batch_size, -1, logits_flat.size(-1)) next_token_log_probs = torch.log_softmax(next_token_logits, dim=-1) - next_token_log_probs, next_token_ids = torch.topk( - next_token_log_probs, self.config.beam_size, dim=-1, largest=True, sorted=True - ) + next_token_log_probs, next_token_ids = torch.topk(next_token_log_probs, + self.config.beam_size, + dim=-1, + largest=True, + sorted=True) # finished sentences is always with EOS, and all but the first one has -inf, so that they will be automatically dropped in the round of beam search. finished_sents = ~input_unfinished_sents next_token_log_probs.masked_fill_(finished_sents.unsqueeze(-1), -numpy.inf) next_token_log_probs[..., 0].masked_fill_(finished_sents, 0) - next_token_ids.masked_fill_( - finished_sents.unsqueeze(-1), self.config.eos_token_id - ) + next_token_ids.masked_fill_(finished_sents.unsqueeze(-1), self.config.eos_token_id) output_log_probs = input_log_probs.unsqueeze(-1) + next_token_log_probs # select N sequences from beams of each input, sorted by sequence probability - output_log_probs = output_log_probs.view( - self.config.batch_size, -1 - ) # shape=(batch, beam_size^2) - output_log_probs, selected_index_flat = output_log_probs.topk( - self.config.beam_size, dim=-1, largest=True, sorted=True - ) # output shape=(batch, beam_size) + output_log_probs = output_log_probs.view(self.config.batch_size, -1) # shape=(batch, beam_size^2) + output_log_probs, selected_index_flat = output_log_probs.topk(self.config.beam_size, + dim=-1, + largest=True, + sorted=True) # output shape=(batch, beam_size) # select the correspondent sentences/next tokens selected_input_seq = selected_index_flat // self.config.beam_size - next_token_ids = next_token_ids.view(self.config.batch_size, -1).gather( - -1, selected_index_flat - ) + next_token_ids = next_token_ids.view(self.config.batch_size, -1).gather(-1, selected_index_flat) - prev_step_results = prev_step_results.view( - self.config.batch_size, -1, prev_step_results.size(-1) - ) + prev_step_results = prev_step_results.view(self.config.batch_size, -1, prev_step_results.size(-1)) prev_step_results = prev_step_results.gather( - 1, selected_input_seq.unsqueeze(-1).repeat(1, 1, prev_step_results.size(-1)) - ) + 1, + selected_input_seq.unsqueeze(-1).repeat(1, 1, prev_step_results.size(-1))) output_unfinished_sents = input_unfinished_sents.gather(1, selected_input_seq) - output_unfinished_sents = ( - output_unfinished_sents - & next_token_ids.ne(self.config.eos_token_id) - ) + output_unfinished_sents = (output_unfinished_sents & next_token_ids.ne(self.config.eos_token_id)) # get the next full input_ids - current_step_results = torch.cat( - [prev_step_results, next_token_ids.unsqueeze(-1)], dim=-1 - ).contiguous() + current_step_results = torch.cat([prev_step_results, next_token_ids.unsqueeze(-1)], dim=-1).contiguous() - prev_step_scores = prev_step_scores.view( - self.config.batch_size, -1, prev_step_scores.size(-1) - ) + prev_step_scores = prev_step_scores.view(self.config.batch_size, -1, prev_step_scores.size(-1)) prev_step_scores = prev_step_scores.gather( - 1, selected_input_seq.unsqueeze(-1).repeat(1, 1, prev_step_scores.size(-1)) - ) - current_step_scores = torch.cat( - [prev_step_scores, output_log_probs.unsqueeze(-1)], dim=-1 - ).contiguous() + 1, + selected_input_seq.unsqueeze(-1).repeat(1, 1, prev_step_scores.size(-1))) + current_step_scores = torch.cat([prev_step_scores, output_log_probs.unsqueeze(-1)], dim=-1).contiguous() return ( next_token_ids, @@ -131,12 +119,192 @@ class GPT2LMHeadModel_BeamSearchStep(GPT2LMHeadModel): ) +class GPT2LMHeadModel_ConfigurableOneStepSearch(GPT2LMHeadModel): + """Here we wrap a class for Onnx model conversion for GPT2LMHeadModel with past state and one + step beam search with configuration support.""" + def __init__(self, + config, + batch_size, + beam_size, + ignore_eos=False, + temperature=1.0, + repetition_penalty=1.0, + excluded_token_ids=None, + length_penalty=1.0, + do_sample=False, + do_sample_top_p=1, + do_sample_top_k=0): + super().__init__(config) + self.config.batch_size = batch_size + self.config.beam_size = beam_size + self.config.ignore_eos = ignore_eos + self.config.temperature = temperature + self.config.repetition_penalty = repetition_penalty + self.config.excluded_token_ids = excluded_token_ids + self.config.length_penalty = length_penalty + self.config.do_sample = do_sample + self.config.do_sample_top_p = do_sample_top_p + self.config.do_sample_top_k = do_sample_top_k + + @staticmethod + def collapse_first_two_dims(tensor): + return tensor.view(-1, *tensor.size()[2:]) + + @staticmethod + def top_k_top_p_filtering(log_probs, top_p=1.0, top_k=0): + '''Set tail event (out of top_p) to a big negative number''' + sorted_log_probs, sorted_indices = torch.sort(log_probs, descending=True) + cumulative_probs = torch.cumsum(sorted_log_probs.exp(), dim=-1) + sorted_indices_to_remove = cumulative_probs >= top_p + sorted_indices_to_remove = torch.cat( + [torch.zeros_like(sorted_indices_to_remove[..., :1]), sorted_indices_to_remove[..., :-1]], dim=-1) + if top_k > 0: + sorted_indices_to_remove = torch.cat( + [sorted_indices_to_remove[..., :top_k], + torch.ones_like(sorted_indices_to_remove[..., top_k:])], dim=-1) + sorted_log_probs.masked_fill_(sorted_indices_to_remove, BIG_NEG) + return log_probs.scatter(-1, sorted_indices, sorted_log_probs) + + def forward( + self, + input_ids, + beam_select_idx, + input_log_probs, + input_unfinished_sents, + prev_step_scores, + *past, + ): + input_ids = input_ids.view(self.config.batch_size, -1, input_ids.size(-1)) + input_num_seq_per_sample = input_ids.size(1) + + input_ids_unfinished_flat = self.collapse_first_two_dims(input_ids).index_select( + 0, + input_unfinished_sents.view(-1).nonzero(as_tuple=False).view(-1)) + + if self.config.ignore_eos: + attention_mask = (input_ids_unfinished_flat != self.config.eos_token_id).float() + else: + attention_mask = torch.ones(input_ids_unfinished_flat.shape).float().to(input_ids_unfinished_flat.device) + position_ids = (attention_mask.cumsum(-1) - 1).clamp(min=0).long() + + if past: + last_seq_len = past[0].size(-2) + input_ids_unfinished_flat = input_ids_unfinished_flat[:, last_seq_len:] + position_ids = position_ids[:, last_seq_len:] + + unfinished_index_relative_to_last_unfinished = beam_select_idx.view(-1)[input_unfinished_sents.view( + -1).nonzero(as_tuple=False).view(-1)] + + past = tuple([p.index_select(1, unfinished_index_relative_to_last_unfinished) for p in past]) + + result = super().forward( + input_ids_unfinished_flat.view(-1, input_ids_unfinished_flat.size(-1)), + position_ids=position_ids, + attention_mask=attention_mask, + past_key_values=past, + return_dict=False, + ) + logits_flat, present_flat = MyGPT2Model.post_process(result, self.config.n_layer) + + # insert finished sequence back to form a square shape of (batch_size, beam_size) + next_token_logits = logits_flat.new_zeros(input_ids.size()[:2] + (logits_flat.size(-1), )) + next_token_logits.index_fill_(2, torch.LongTensor([self.config.eos_token_id]).to(input_ids.device), -BIG_NEG) + + next_token_logits.masked_scatter_( + input_unfinished_sents.unsqueeze(-1).expand_as(next_token_logits), logits_flat[:, -1]) + + # repetition penalty from CTRL paper (https://arxiv.org/abs/1909.05858) + if self.config.repetition_penalty != 1.0: + _pen = next_token_logits.gather(2, input_ids) + _pen = torch.where(_pen > 0, _pen / self.config.repetition_penalty, _pen * self.config.repetition_penalty) + next_token_logits.scatter_(2, input_ids, _pen) + + # similar way to encourage short sentence + if self.config.length_penalty != 1.0: + _pen = next_token_logits[..., self.config.eos_token_id] + # if eos > 0, increase it, else, decrease it. + _pen = torch.where(_pen > 0, _pen * self.config.length_penalty, _pen / self.config.length_penalty) + next_token_logits[..., self.config.eos_token_id] = _pen + + if self.config.temperature != 1.0: + next_token_logits = next_token_logits / self.config.temperature + + # exclude excluded_token_ids + if self.config.excluded_token_ids is not None: + next_token_logits.index_fill_(2, self.config.excluded_token_ids.to(next_token_logits.device), + BIG_NEG) # batch x beams/sequences x vocab_size + + next_token_log_probs = torch.log_softmax(next_token_logits, dim=-1) + + if self.config.do_sample: + vocab_size = next_token_log_probs.size(-1) + _next_token_log_probs = self.top_k_top_p_filtering(next_token_log_probs.view(-1, vocab_size), + top_k=self.config.do_sample_top_k, + top_p=self.config.do_sample_top_p) + next_token_ids = torch.multinomial(_next_token_log_probs.exp(), + num_samples=self.config.beam_size, + replacement=False) + next_token_ids = next_token_ids.view(self.config.batch_size, input_num_seq_per_sample, -1) + next_token_log_probs = next_token_log_probs.gather(-1, next_token_ids) + else: + next_token_log_probs, next_token_ids = torch.topk(next_token_log_probs, + self.config.beam_size, + dim=-1, + largest=True, + sorted=True) + + output_log_probs = input_log_probs.unsqueeze(-1) + next_token_log_probs + + # select N sequences from beams of each input, sorted by sequence probability + output_log_probs = output_log_probs.view(self.config.batch_size, -1) # shape=(batch, beam_size^2) + output_log_probs, selected_index_flat = output_log_probs.topk(self.config.beam_size, + dim=-1, + largest=True, + sorted=True) # output shape=(batch, beam_size) + + # select the correspondent sentences/next tokens + selected_input_seq = selected_index_flat // self.config.beam_size + next_token_ids = next_token_ids.view(self.config.batch_size, -1).gather(-1, selected_index_flat) + + prev_step_results = input_ids.view(self.config.batch_size, -1, input_ids.size(-1)).contiguous() + prev_step_results = prev_step_results.gather( + 1, + selected_input_seq.unsqueeze(-1).expand(selected_input_seq.shape + (prev_step_results.size(-1), ))) + + output_unfinished_sents = input_unfinished_sents.gather(1, selected_input_seq) + output_unfinished_sents = (output_unfinished_sents & next_token_ids.ne(self.config.eos_token_id)) + + current_step_results = torch.cat([prev_step_results, next_token_ids.unsqueeze(-1)], dim=-1).contiguous() + + prev_step_scores = prev_step_scores.view(self.config.batch_size, -1, prev_step_scores.size(-1)) + prev_step_scores = prev_step_scores.gather( + 1, + selected_input_seq.unsqueeze(-1).expand(selected_input_seq.shape + (prev_step_scores.size(-1), ))) + current_step_scores = torch.cat([prev_step_scores, output_log_probs.unsqueeze(-1)], dim=-1).contiguous() + + # For next past state + index_relative_to_last_unfinished = (input_unfinished_sents.view(-1).float().cumsum(-1) - 1).clamp( + min=0).long().reshape_as(input_unfinished_sents).gather(1, selected_input_seq) + + return ( + current_step_results.view(self.config.batch_size * self.config.beam_size, -1), + present_flat, + index_relative_to_last_unfinished, + output_log_probs, + output_unfinished_sents, + current_step_scores.view(self.config.batch_size * self.config.beam_size, -1), + ) + + # Maps model class name to a tuple of model class, name of first output and use padding or not MODEL_CLASSES = { 'GPT2LMHeadModel': (MyGPT2LMHeadModel, 'logits', True), 'GPT2LMHeadModel_NoPadding': (MyGPT2LMHeadModel_NoPadding, 'logits', False), 'GPT2Model': (MyGPT2Model, 'last_state', True), - "GPT2LMHeadModel_BeamSearchStep": (GPT2LMHeadModel_BeamSearchStep, "last_state", True), # defined in gpt2_beamsearch_helper.py + "GPT2LMHeadModel_BeamSearchStep": + (GPT2LMHeadModel_BeamSearchStep, "last_state", True), # defined in gpt2_beamsearch_helper.py + "GPT2LMHeadModel_ConfigurableOneStepSearch": + (GPT2LMHeadModel_ConfigurableOneStepSearch, "last_state", False), # defined in gpt2_beamsearch_helper.py } @@ -144,22 +312,20 @@ class Gpt2BeamSearchInputs(Gpt2Inputs): def __init__( self, input_ids, + past, position_ids, attention_mask, - past, beam_select_idx=None, input_log_probs=None, input_unfinished_sents=None, prev_step_results=None, prev_step_scores=None, ): - super().__init__(input_ids, position_ids, attention_mask, past) + super().__init__(input_ids, position_ids, attention_mask, past=past) self.prev_step_results: torch.LongTensor = prev_step_results self.prev_step_scores: Union[torch.FloatTensor, torch.HalfTensor, torch.cuda.FloatTensor] = prev_step_scores if beam_select_idx is None: - self.beam_select_idx: torch.LongTensor = torch.zeros( - [1, len(input_ids)] - ).long() + self.beam_select_idx: torch.LongTensor = torch.zeros([1, len(input_ids)]).long() else: self.beam_select_idx: torch.LongTensor = beam_select_idx self.input_log_probs: Union[torch.FloatTensor, torch.HalfTensor, torch.cuda.FloatTensor] = input_log_probs @@ -167,30 +333,24 @@ class Gpt2BeamSearchInputs(Gpt2Inputs): def to_list(self) -> List: input_list = [ - v - for v in [ - self.input_ids, - self.position_ids, - self.attention_mask, - self.beam_select_idx, - self.input_log_probs, - self.input_unfinished_sents, - self.prev_step_results, - self.prev_step_scores - ] - if v is not None + v for v in [ + self.input_ids, self.position_ids, self.attention_mask, self.beam_select_idx, self.input_log_probs, + self.input_unfinished_sents, self.prev_step_results, self.prev_step_scores + ] if v is not None ] if self.past: input_list.extend(self.past) return input_list def to_fp32(self): - gpt2_inputs = super().to_fp32() + past = [p.to(dtype=torch.float32) for p in self.past] + attention_mask = self.attention_mask.to( + dtype=torch.float32) if self.attention_mask is not None else self.attention_mask return Gpt2BeamSearchInputs( - gpt2_inputs.input_ids, - gpt2_inputs.position_ids, - gpt2_inputs.attention_mask, - gpt2_inputs.past, + self.input_ids, + past, + self.position_ids, + attention_mask, self.beam_select_idx, self.input_log_probs.to(dtype=torch.float32), self.input_unfinished_sents, @@ -201,7 +361,6 @@ class Gpt2BeamSearchInputs(Gpt2Inputs): class Gpt2BeamSearchHelper(Gpt2Helper): """A helper class for Gpt2 model conversion, inference and verification.""" - @staticmethod def get_dummy_inputs(batch_size: int, past_sequence_length: int, @@ -217,40 +376,32 @@ class Gpt2BeamSearchHelper(Gpt2Helper): """Create random inputs for GPT2 model. Returns torch tensors of input_ids, position_ids, attention_mask and a list of past state tensors. """ - gpt2_dummy_inputs = Gpt2Helper.get_dummy_inputs( - batch_size, - past_sequence_length, - sequence_length, - num_attention_heads, - hidden_size, - num_layer, - vocab_size, - device, - float16, - has_position_ids, - has_attention_mask - ) + gpt2_dummy_inputs = Gpt2Helper.get_dummy_inputs(batch_size, past_sequence_length, sequence_length, + num_attention_heads, hidden_size, num_layer, vocab_size, device, + float16, has_position_ids, has_attention_mask) float_type = torch.float16 if float16 else torch.float32 beam_select_idx = torch.zeros([1, batch_size], device=device).long() input_log_probs = torch.zeros([batch_size, 1], dtype=float_type, device=device) - input_unfinished_sents = torch.ones( - [batch_size, 1], dtype=torch.bool, device=device - ) - prev_step_results = torch.randint( - low=0, - high=vocab_size - 1, - size=(batch_size, sequence_length), - dtype=torch.int64, - device=device, - ) + input_unfinished_sents = torch.ones([batch_size, 1], dtype=torch.bool, device=device) + if has_position_ids: + prev_step_results = torch.randint( + low=0, + high=vocab_size - 1, + size=(batch_size, sequence_length), + dtype=torch.int64, + device=device, + ) + else: + prev_step_results = None + prev_step_scores = torch.zeros([batch_size, 1], dtype=float_type, device=device) return Gpt2BeamSearchInputs( gpt2_dummy_inputs.input_ids, + gpt2_dummy_inputs.past, gpt2_dummy_inputs.position_ids, gpt2_dummy_inputs.attention_mask, - gpt2_dummy_inputs.past, beam_select_idx, input_log_probs, input_unfinished_sents, @@ -266,16 +417,21 @@ class Gpt2BeamSearchHelper(Gpt2Helper): beam_size: int, step: int, config: GPT2Config, - model_class: str = "GPT2LMHeadModel") -> Dict[str, List[int]]: + model_class: str = "GPT2LMHeadModel", + num_seq: int = 0) -> Dict[str, List[int]]: """Returns a dictionary with output name as key, and shape as value.""" num_attention_heads = config.num_attention_heads hidden_size = config.hidden_size num_layer = config.num_hidden_layers - vocab_size = config.vocab_size + vocab_size = config.vocab_size output_name = MODEL_CLASSES[model_class][1] - last_state_shape = [batch_size, beam_size] + if model_class == "GPT2LMHeadModel_BeamSearchStep": + last_state_shape = [batch_size, beam_size] + else: + last_state_shape = [batch_size * beam_size, past_sequence_length + sequence_length + 1] + if step == 0: present_state_shape = [ 2, @@ -285,9 +441,12 @@ class Gpt2BeamSearchHelper(Gpt2Helper): int(hidden_size / num_attention_heads), ] else: + if num_seq == 0: + num_seq = beam_size + present_state_shape = [ 2, - batch_size * beam_size, + batch_size * num_seq, num_attention_heads, past_sequence_length + sequence_length, int(hidden_size / num_attention_heads), @@ -300,49 +459,49 @@ class Gpt2BeamSearchHelper(Gpt2Helper): output_shapes["output_selected_indices"] = [1, batch_size * beam_size] output_shapes["output_log_probs"] = [batch_size, beam_size] output_shapes["output_unfinished_sents"] = [batch_size, beam_size] - output_shapes["current_step_results"] = [batch_size * beam_size, past_sequence_length + sequence_length + 1] - output_shapes["current_step_scores"] = [batch_size * beam_size, past_sequence_length + sequence_length - context_len + 2] + if model_class == "GPT2LMHeadModel_BeamSearchStep": + output_shapes["current_step_results"] = [batch_size * beam_size, past_sequence_length + sequence_length + 1] + output_shapes["current_step_scores"] = [ + batch_size * beam_size, past_sequence_length + sequence_length - context_len + 2 + ] return output_shapes @staticmethod - def get_output_buffers( - output_shapes, device, is_float16=False - ): + def get_output_buffers(output_shapes, device, is_float16=False): """Returns a dictionary of output name as key, and 1D tensor as value. The tensor has enough space for given shape.""" data_type = torch.float16 if is_float16 else torch.float32 output_buffers = {} for name, shape in output_shapes.items(): - if ( - name == "output_selected_indices" - or name == "current_step_results" - or name == "last_state" - ): - output_buffers[name] = torch.empty( - numpy.prod(shape), dtype=torch.long, device=device - ) + if (name == "output_selected_indices" or name == "current_step_results" or name == "last_state"): + output_buffers[name] = torch.empty(numpy.prod(shape), dtype=torch.long, device=device) elif name == "output_unfinished_sents": - output_buffers[name] = torch.empty( - numpy.prod(shape), dtype=torch.bool, device=device - ) + output_buffers[name] = torch.empty(numpy.prod(shape), dtype=torch.bool, device=device) else: - output_buffers[name] = torch.empty( - numpy.prod(shape), dtype=data_type, device=device - ) + output_buffers[name] = torch.empty(numpy.prod(shape), dtype=data_type, device=device) return output_buffers @staticmethod - def compare_outputs(torch_outputs, ort_outputs, rtol=1e-03, atol=1e-03): + def compare_outputs(torch_outputs, + ort_outputs, + model_class="GPT2LMHeadModel_BeamSearchStep", + rtol=1e-03, + atol=1e-03): """Returns True if torch and ORT outputs are close for given thresholds, and False otherwise.""" - is_close = numpy.allclose( - ort_outputs[-4], torch_outputs[-4].cpu().numpy(), rtol=rtol, atol=atol - ) - logger.debug( - f"PyTorch and OnnxRuntime output 0 (last_state) are close: {is_close}" - ) + if model_class == "GPT2LMHeadModel_BeamSearchStep": + results_id = -4 + num_layers = len(ort_outputs) - 6 + else: + results_id = 0 + num_layers = len(ort_outputs) - 5 + + is_close = numpy.allclose(ort_outputs[results_id], + torch_outputs[results_id].cpu().numpy(), + rtol=rtol, + atol=atol) + logger.debug(f"PyTorch and OnnxRuntime output 0 (last_state) are close: {is_close}") is_all_close = is_close - num_layers = len(ort_outputs) - 6 for layer in range(num_layers): is_close = numpy.allclose( ort_outputs[1 + layer], @@ -350,16 +509,12 @@ class Gpt2BeamSearchHelper(Gpt2Helper): rtol=rtol, atol=atol, ) - logger.debug( - f"PyTorch and OnnxRuntime layer {layer} state (present_{layer}) are close:{is_close}" - ) + logger.debug(f"PyTorch and OnnxRuntime layer {layer} state (present_{layer}) are close:{is_close}") is_all_close = is_all_close and is_close if not is_all_close: max_abs_diff = Gpt2BeamSearchHelper.diff_outputs(torch_outputs, ort_outputs) - logger.info( - f"PyTorch and OnnxRuntime results are not all close: max_abs_diff={max_abs_diff:.5f}" - ) + logger.info(f"PyTorch and OnnxRuntime results are not all close: max_abs_diff={max_abs_diff:.5f}") return is_all_close @@ -376,7 +531,7 @@ class Gpt2BeamSearchHelper(Gpt2Helper): num_layer = config.n_layer dummy_inputs = Gpt2BeamSearchHelper.get_dummy_inputs(batch_size=1, past_sequence_length=1, - sequence_length=1, + sequence_length=2, num_attention_heads=config.num_attention_heads, hidden_size=config.hidden_size, num_layer=num_layer, @@ -396,13 +551,21 @@ class Gpt2BeamSearchHelper(Gpt2Helper): output_names = ["last_state"] + present_names - output_names += [ - "output_selected_indices", - "output_log_probs", - "output_unfinished_sents", - "current_step_results", - "current_step_scores", - ] + if has_position_ids: + output_names += [ + "output_selected_indices", + "output_log_probs", + "output_unfinished_sents", + "current_step_results", + "current_step_scores", + ] + else: + output_names += [ + "output_selected_indices", + "output_log_probs", + "output_unfinished_sents", + "current_step_scores", + ] # Shape of input tensors: # input_ids: (batch_size, seq_len) @@ -413,27 +576,36 @@ class Gpt2BeamSearchHelper(Gpt2Helper): # or logits: (batch_size, seq_len, vocab_size) # present_{i}: (2, batch_size, num_heads, past_seq_len + seq_len, hidden_size/num_heads) dynamic_axes = { - "input_ids": {0: "batch_size", 1: "seq_len"}, - output_names[0]: {0: "batch_size", 1: "seq_len"}, + "input_ids": { + 0: "batch_size", + 1: "seq_len" + }, + output_names[0]: { + 0: "batch_size", + 1: "seq_len" + }, } for name in past_names: dynamic_axes[name] = {1: "batch_size", 3: "past_seq_len"} for name in present_names: - dynamic_axes[name] = {1: "batch_size", 3: "total_seq_len"} + dynamic_axes[name] = {1: "batch_size", 3: "cur_seq_len"} input_names = ["input_ids"] - dynamic_axes["position_ids"] = {0: "batch_size", 1: "seq_len"} - input_names.append("position_ids") - dynamic_axes["attention_mask"] = {0: "batch_size", 1: "total_seq_len"} - input_names.append("attention_mask") + if has_position_ids: + dynamic_axes["position_ids"] = {0: "batch_size", 1: "seq_len"} + input_names.append("position_ids") + if has_attention_mask: + dynamic_axes["attention_mask"] = {0: "batch_size", 1: "total_seq_len"} + input_names.append("attention_mask") dynamic_axes["beam_select_idx"] = {1: "batch_size"} input_names.append("beam_select_idx") dynamic_axes["input_log_probs"] = {0: "batch_size", 1: "beam_size"} input_names.append("input_log_probs") dynamic_axes["input_unfinished_sents"] = {0: "batch_size", 1: "beam_size"} input_names.append("input_unfinished_sents") - dynamic_axes["prev_step_results"] = {0: "batch_size", 1: "total_seq_len"} - input_names.append("prev_step_results") + if has_position_ids: + dynamic_axes["prev_step_results"] = {0: "batch_size", 1: "total_seq_len"} + input_names.append("prev_step_results") dynamic_axes["prev_step_scores"] = {0: "batch_size", 1: "total_seq_len"} input_names.append("prev_step_scores") input_names.extend(past_names) @@ -463,38 +635,22 @@ class Gpt2BeamSearchHelper(Gpt2Helper): """Run inference of ONNX model, and returns average latency in ms when total_runs > 0 besides outputs.""" logger.debug(f"start onnxruntime_inference") - ort_inputs = { - "input_ids": numpy.ascontiguousarray(inputs.input_ids.cpu().numpy()) - } + ort_inputs = {"input_ids": numpy.ascontiguousarray(inputs.input_ids.cpu().numpy())} if inputs.position_ids is not None: - ort_inputs["position_ids"] = numpy.ascontiguousarray( - inputs.position_ids.cpu().numpy() - ) + ort_inputs["position_ids"] = numpy.ascontiguousarray(inputs.position_ids.cpu().numpy()) if inputs.attention_mask is not None: - ort_inputs["attention_mask"] = numpy.ascontiguousarray( - inputs.attention_mask.cpu().numpy() - ) + ort_inputs["attention_mask"] = numpy.ascontiguousarray(inputs.attention_mask.cpu().numpy()) if inputs.beam_select_idx is not None: - ort_inputs["beam_select_idx"] = numpy.ascontiguousarray( - inputs.beam_select_idx.cpu().numpy() - ) + ort_inputs["beam_select_idx"] = numpy.ascontiguousarray(inputs.beam_select_idx.cpu().numpy()) if inputs.input_log_probs is not None: - ort_inputs["input_log_probs"] = numpy.ascontiguousarray( - inputs.input_log_probs.cpu().numpy() - ) + ort_inputs["input_log_probs"] = numpy.ascontiguousarray(inputs.input_log_probs.cpu().numpy()) if inputs.input_unfinished_sents is not None: - ort_inputs["input_unfinished_sents"] = numpy.ascontiguousarray( - inputs.input_unfinished_sents.cpu().numpy() - ) + ort_inputs["input_unfinished_sents"] = numpy.ascontiguousarray(inputs.input_unfinished_sents.cpu().numpy()) if inputs.prev_step_results is not None: - ort_inputs["prev_step_results"] = numpy.ascontiguousarray( - inputs.prev_step_results.cpu().numpy() - ) + ort_inputs["prev_step_results"] = numpy.ascontiguousarray(inputs.prev_step_results.cpu().numpy()) if inputs.prev_step_scores is not None: - ort_inputs["prev_step_scores"] = numpy.ascontiguousarray( - inputs.prev_step_scores.cpu().numpy() - ) + ort_inputs["prev_step_scores"] = numpy.ascontiguousarray(inputs.prev_step_scores.cpu().numpy()) if inputs.past is not None: for i, past_i in enumerate(inputs.past): ort_inputs[f"past_{i}"] = numpy.ascontiguousarray(past_i.cpu().numpy()) @@ -510,9 +666,7 @@ class Gpt2BeamSearchHelper(Gpt2Helper): latency.append(time.time() - start) average_latency = sum(latency) * 1000 / len(latency) - logger.debug( - "OnnxRuntime Inference time = {} ms".format(format(average_latency, ".2f")) - ) + logger.debug("OnnxRuntime Inference time = {} ms".format(format(average_latency, ".2f"))) return ort_outputs, average_latency @@ -532,11 +686,18 @@ class Gpt2BeamSearchHelper(Gpt2Helper): """Returnas IO binding object for a session.""" # Bind inputs and outputs to onnxruntime session - io_binding = Gpt2Helper.prepare_io_binding(ort_session, input_ids, position_ids, attention_mask, past, output_buffers, output_shapes) + io_binding = Gpt2Helper.prepare_io_binding(ort_session, + input_ids, + position_ids, + attention_mask, + past=past, + output_buffers=output_buffers, + output_shapes=output_shapes) # Bind inputs data_type = output_buffers[ort_session.get_outputs()[1].name].dtype float_type = numpy.float16 if data_type == torch.float16 else numpy.float32 + if past is not None: for i, past_i in enumerate(past): assert past_i.is_contiguous() @@ -613,14 +774,9 @@ class Gpt2BeamSearchHelper(Gpt2Helper): for output in ort_session.get_outputs(): output_name = output.name output_buffer = output_buffers[output_name] - logger.debug( - f"{output_name} device type={output_buffer.device.type} shape={list(output_buffer.size())}" - ) - if ( - output_name == "output_selected_indices" - or output_name == "last_state" - or output_name == "current_step_results" - ): + logger.debug(f"{output_name} device type={output_buffer.device.type} shape={list(output_buffer.size())}") + if (output_name == "output_selected_indices" or output_name == "last_state" + or output_name == "current_step_results"): io_binding.bind_output( output_name, output_buffer.device.type, @@ -682,9 +838,8 @@ class Gpt2BeamSearchHelper(Gpt2Helper): ort_session.run_with_iobinding(io_binding) # Copy results to cpu for verification - ort_outputs = Gpt2BeamSearchHelper.get_outputs_from_io_binding_buffer( - ort_session, output_buffers, output_shapes, return_numpy - ) + ort_outputs = Gpt2BeamSearchHelper.get_outputs_from_io_binding_buffer(ort_session, output_buffers, + output_shapes, return_numpy) if total_runs == 0: return ort_outputs @@ -695,17 +850,12 @@ class Gpt2BeamSearchHelper(Gpt2Helper): # Run onnxruntime with io binding ort_session.run_with_iobinding(io_binding) if include_copy_output_latency: - _ = Gpt2BeamSearchHelper.get_outputs_from_io_binding_buffer( - ort_session, output_buffers, output_shapes, return_numpy - ) + _ = Gpt2BeamSearchHelper.get_outputs_from_io_binding_buffer(ort_session, output_buffers, output_shapes, + return_numpy) latency.append(time.time() - start) average_latency = sum(latency) * 1000 / len(latency) - logger.debug( - "OnnxRuntime with IO binding inference time = {} ms".format( - format(average_latency, ".2f") - ) - ) + logger.debug("OnnxRuntime with IO binding inference time = {} ms".format(format(average_latency, ".2f"))) return ort_outputs, average_latency @@ -746,38 +896,24 @@ class Gpt2BeamSearchHelper(Gpt2Helper): config, model_class, ) - output_buffers = Gpt2BeamSearchHelper.get_output_buffers( - max_output_shapes, device, is_float16 - ) + output_buffers = Gpt2BeamSearchHelper.get_output_buffers(max_output_shapes, device, is_float16) passed_test_cases = 0 for _ in range(total_test_cases): - sequence_length = random.randint(1, max_seq_len) past_sequence_length = random.randint(0, max_past_seq_len) + sequence_length = random.randint(1 + past_sequence_length, max_seq_len + past_sequence_length) batch_size = random.randint(1, max_batch_size) logger.debug( - f"Running parity test for batch_size={batch_size} past_sequence_length={past_sequence_length}..." - ) - dummy_inputs = Gpt2BeamSearchHelper.get_dummy_inputs( - batch_size, - past_sequence_length, - sequence_length, - config.num_attention_heads, - config.hidden_size, - config.n_layer, - config.vocab_size, - device, - is_float16, - has_position_ids, - has_attention_mask - ) + f"Running parity test for batch_size={batch_size} past_sequence_length={past_sequence_length}...") + dummy_inputs = Gpt2BeamSearchHelper.get_dummy_inputs(batch_size, past_sequence_length, sequence_length, + config.num_attention_heads, config.hidden_size, + config.n_layer, config.vocab_size, device, is_float16, + has_position_ids, has_attention_mask) outputs = Gpt2BeamSearchHelper.pytorch_inference(model, dummy_inputs) if use_io_binding: - ort_outputs = Gpt2BeamSearchHelper.onnxruntime_inference( - ort_session, dummy_inputs - ) + ort_outputs = Gpt2BeamSearchHelper.onnxruntime_inference(ort_session, dummy_inputs) else: output_shapes = Gpt2BeamSearchHelper.get_output_shapes( batch_size, @@ -790,19 +926,18 @@ class Gpt2BeamSearchHelper(Gpt2Helper): model_class, ) ort_outputs = Gpt2BeamSearchHelper.onnxruntime_inference_with_binded_io( - ort_session, dummy_inputs, output_buffers, output_shapes - ) + ort_session, dummy_inputs, output_buffers, output_shapes) - is_all_close = Gpt2BeamSearchHelper.compare_outputs( - outputs, ort_outputs, rtol=rtol, atol=atol - ) + is_all_close = Gpt2BeamSearchHelper.compare_outputs(outputs, + ort_outputs, + model_class=model_class, + rtol=rtol, + atol=atol) if is_all_close: passed_test_cases += 1 logger.info(f"Parity Test Cases={total_test_cases}; Passed={passed_test_cases}") if passed_test_cases > 0.95 * total_test_cases: - logger.info( - f"Parity is good: passed rate={int(passed_test_cases*100/total_test_cases):.0f}%" - ) + logger.info(f"Parity is good: passed rate={int(passed_test_cases*100/total_test_cases):.0f}%") return passed_test_cases == total_test_cases @staticmethod diff --git a/onnxruntime/python/tools/transformers/gpt2_beamsearch_tester.py b/onnxruntime/python/tools/transformers/gpt2_beamsearch_tester.py index 5f0c22b97e..a8f7047bde 100644 --- a/onnxruntime/python/tools/transformers/gpt2_beamsearch_tester.py +++ b/onnxruntime/python/tools/transformers/gpt2_beamsearch_tester.py @@ -20,47 +20,49 @@ from benchmark_helper import Precision logger = logging.getLogger(__name__) + class Gpt2TesterFactory: @staticmethod def create_tester(tester_type="default"): testers = { "default": Gpt2Tester, "beam_search_step": Gpt2BeamSearchTester, + "configurable_one_step_search": Gpt2BeamSearchTester, } w = testers[tester_type] return w + class Gpt2BeamSearchTester(Gpt2Tester): - def __init__(self, - input_ids, - position_ids, - attention_mask, - beam_select_idx, - input_log_probs, - input_unfinished_sents, - prev_step_results, - prev_step_scores, - num_attention_heads, - hidden_size, - num_layer, - beam_size, - device, - is_fp16=False, - top_k=20, - top_k_required_order=False, + def __init__( + self, + input_ids, + position_ids, + attention_mask, + beam_select_idx, + input_log_probs, + input_unfinished_sents, + prev_step_results, + prev_step_scores, + num_attention_heads, + hidden_size, + num_layer, + beam_size, + device, + is_fp16=False, + top_k=20, + top_k_required_order=False, ): - super().__init__( - input_ids, - position_ids, - attention_mask, - num_attention_heads, - hidden_size, - num_layer, - device, - is_fp16, - top_k, - top_k_required_order - ) + super().__init__(input_ids, + position_ids, + attention_mask, + num_attention_heads=num_attention_heads, + hidden_size=hidden_size, + num_layer=num_layer, + device=device, + is_fp16=is_fp16, + top_k=top_k, + top_k_required_order=top_k_required_order) self.input_length = input_ids.shape[-1] self.n_layer = num_layer self.beam_size = beam_size @@ -71,7 +73,7 @@ class Gpt2BeamSearchTester(Gpt2Tester): self.input_log_probs = input_log_probs.type(float_type).to(device) self.input_unfinished_sents = input_unfinished_sents.to(device) - self.prev_step_results = prev_step_results.to(device) + self.prev_step_results = prev_step_results.to(device) if prev_step_results is not None else None self.prev_step_scores = prev_step_scores.type(float_type).to(device) self.last_state = None @@ -79,9 +81,9 @@ class Gpt2BeamSearchTester(Gpt2Tester): def get_inputs(self) -> Gpt2BeamSearchInputs: return Gpt2BeamSearchInputs( self.input_ids, + self.past, self.position_ids, self.attention_mask, - self.past, self.beam_select_idx, self.input_log_probs, self.input_unfinished_sents, @@ -93,74 +95,52 @@ class Gpt2BeamSearchTester(Gpt2Tester): """ Update the inputs for next inference. """ - self.last_state = ( - torch.from_numpy(output[0]).to(device) - if isinstance(output[0], numpy.ndarray) - else output[0].clone().detach().cpu() - ) + self.last_state = (torch.from_numpy(output[0]).to(device) + if isinstance(output[0], numpy.ndarray) else output[0].clone().detach().cpu()) self.input_ids = self.last_state.view(self.batch_size * self.beam_size, -1).to(device) - self.beam_select_idx = ( - torch.from_numpy(output[-5]).to(device) - if isinstance(output[-5], numpy.ndarray) - else output[-5].clone().detach().to(device) - ) - self.input_log_probs = ( - torch.from_numpy(output[-4]).to(device) - if isinstance(output[-4], numpy.ndarray) - else output[-4].clone().detach().to(device) - ) - self.input_unfinished_sents = ( - torch.from_numpy(output[-3]).to(device) - if isinstance(output[-3], numpy.ndarray) - else output[-3].clone().detach().to(device) - ) - self.prev_step_results = ( - torch.from_numpy(output[-2]).to(device) - if isinstance(output[-2], numpy.ndarray) - else output[-2].clone().detach().to(device) - ) - self.prev_step_scores = ( - torch.from_numpy(output[-1]).to(device) - if isinstance(output[-1], numpy.ndarray) - else output[-1].clone().detach().to(device) - ) + if self.position_ids is not None: + input_unfinished_sents_id = -3 + self.prev_step_results = (torch.from_numpy(output[-2]).to(device) if isinstance(output[-2], numpy.ndarray) + else output[-2].clone().detach().to(device)) + self.position_ids = (torch.tensor([self.input_length + step - 1 + ]).unsqueeze(0).repeat(self.batch_size * self.beam_size, 1).to(device)) + + if self.attention_mask.size(0) != (self.batch_size * self.beam_size): + self.attention_mask = self.attention_mask.repeat(self.batch_size * self.beam_size, 1) + self.attention_mask = torch.cat( + [ + self.attention_mask, + torch.ones([self.batch_size * self.beam_size, 1]).type_as(self.attention_mask), + ], + 1, + ).to(device) + else: + input_unfinished_sents_id = -2 + + self.beam_select_idx = (torch.from_numpy(output[input_unfinished_sents_id - 2]).to(device) if isinstance( + output[input_unfinished_sents_id - 2], numpy.ndarray) else output[input_unfinished_sents_id - + 2].clone().detach().to(device)) + self.input_log_probs = (torch.from_numpy(output[input_unfinished_sents_id - 1]).to(device) if isinstance( + output[input_unfinished_sents_id - 1], numpy.ndarray) else output[input_unfinished_sents_id - + 1].clone().detach().to(device)) + self.input_unfinished_sents = (torch.from_numpy(output[input_unfinished_sents_id]).to(device) if isinstance( + output[input_unfinished_sents_id], numpy.ndarray) else + output[input_unfinished_sents_id].clone().detach().to(device)) + self.prev_step_scores = (torch.from_numpy(output[-1]).to(device) + if isinstance(output[-1], numpy.ndarray) else output[-1].clone().detach().to(device)) self.top_1_tokens = self.input_ids[0] self.top_k_tokens = self.last_state - self.position_ids = ( - torch.tensor([self.input_length + step - 1]) - .unsqueeze(0) - .repeat(self.batch_size * self.beam_size, 1) - .to(device) - ) - - if self.attention_mask.size(0) != (self.batch_size * self.beam_size): - self.attention_mask = self.attention_mask.repeat( - self.batch_size * self.beam_size, 1 - ) - self.attention_mask = torch.cat( - [ - self.attention_mask, - torch.ones([self.batch_size * self.beam_size, 1]).type_as( - self.attention_mask - ), - ], - 1, - ).to(device) - self.past = [] if isinstance(output[1], tuple): # past in torch output is tuple self.past = list(output[1]) else: for i in range(self.n_layer): - past_i = ( - torch.from_numpy(output[i + 1]) - if isinstance(output[i + 1], numpy.ndarray) - else output[i + 1].clone().detach() - ) + past_i = (torch.from_numpy(output[i + 1]) + if isinstance(output[i + 1], numpy.ndarray) else output[i + 1].clone().detach()) self.past.append(past_i.to(device)) @staticmethod @@ -217,9 +197,7 @@ class Gpt2BeamSearchTester(Gpt2Tester): treatment_name = "Quantized Onnx" if precision == Precision.INT8 else "Onnx" torch_metric = Gpt2Metric(baseline_name, baseline_name, top_k) onnx_metric = Gpt2Metric(treatment_name, baseline_name, top_k) - onnx_io_metric = Gpt2Metric( - treatment_name + " with IO Binding", baseline_name, top_k - ) + onnx_io_metric = Gpt2Metric(treatment_name + " with IO Binding", baseline_name, top_k) for i, inputs in enumerate(test_inputs): if max_inputs > 0 and i == max_inputs: @@ -228,17 +206,15 @@ class Gpt2BeamSearchTester(Gpt2Tester): print(f"{i}") input_ids = inputs["input_ids"] position_ids = inputs["position_ids"] if "position_ids" in inputs else None - attention_mask = ( - inputs["attention_mask"] if "attention_mask" in inputs else None - ) - beam_select_idx = ( - inputs["beam_select_idx"] if "beam_select_idx" in inputs else None - ) - input_log_probs = ( - inputs["input_log_probs"] if "input_log_probs" in inputs else None - ) + attention_mask = (inputs["attention_mask"] if "attention_mask" in inputs else None) + beam_select_idx = (inputs["beam_select_idx"] if "beam_select_idx" in inputs else None) + input_log_probs = (inputs["input_log_probs"] if "input_log_probs" in inputs else None) input_unfinished_sents = inputs["input_unfinished_sents"] - prev_step_results = inputs["input_ids"] + if model_class == "GPT2LMHeadModel_BeamSearchStep": + prev_step_results = inputs["input_ids"] + else: + prev_step_results = None + if "prev_step_scores" in inputs: prev_step_scores = inputs["prev_step_scores"] else: @@ -302,32 +278,36 @@ class Gpt2BeamSearchTester(Gpt2Tester): batch_size = torch_runner.batch_size onnx_metric.start_batch(batch_size) onnx_io_metric.start_batch(batch_size) - context_len = list(onnx_runner.input_ids.size())[1] + context_len = list(onnx_runner.input_ids.size())[-1] with torch.no_grad(): - done = torch.zeros(batch_size, dtype=torch.bool) for step in range(max_steps): print(f"Processing step: {step}") - seq_len = list(onnx_runner.input_ids.size())[1] - past_seq_len = list(onnx_runner.past[0].size())[3] + if model_class == "GPT2LMHeadModel_BeamSearchStep": + num_seq = beam_size + seq_len = list(onnx_runner.input_ids.size())[1] + past_seq_len = list(onnx_runner.past[0].size())[3] + else: + num_seq = sum(onnx_io_runner.input_unfinished_sents.view(-1).long().cpu()) + past_seq_len = list(onnx_runner.past[0].size())[3] + seq_len = list(onnx_runner.input_ids.size())[-1] - past_seq_len start_time = timeit.default_timer() - pytorch_output = Gpt2BeamSearchHelper.pytorch_inference( - model, torch_runner.get_inputs() - ) - torch_metric.add_latency( - past_seq_len, timeit.default_timer() - start_time - ) + pytorch_output = Gpt2BeamSearchHelper.pytorch_inference(model, torch_runner.get_inputs()) + torch_metric.add_latency(past_seq_len, timeit.default_timer() - start_time) torch_runner.update(pytorch_output, step, device) ( onnx_output, avg_latency_ms, - ) = Gpt2BeamSearchHelper.onnxruntime_inference( - session, onnx_runner.get_inputs(), total_runs=1 - ) + ) = Gpt2BeamSearchHelper.onnxruntime_inference(session, onnx_runner.get_inputs(), total_runs=1) onnx_metric.add_latency(past_seq_len, avg_latency_ms / 1000.0) onnx_runner.update(onnx_output, step, device) + if model_class == "GPT2LMHeadModel_BeamSearchStep": + num_seq = beam_size + else: + num_seq = sum(onnx_io_runner.input_unfinished_sents.view(-1).long().cpu()) + output_shapes = Gpt2BeamSearchHelper.get_output_shapes( batch_size, context_len, @@ -336,12 +316,11 @@ class Gpt2BeamSearchTester(Gpt2Tester): beam_size, step, model.config, - model_class=model_class + model_class=model_class, + num_seq=num_seq, ) - Gpt2BeamSearchHelper.auto_increase_buffer_size( - output_buffers, output_shapes - ) + Gpt2BeamSearchHelper.auto_increase_buffer_size(output_buffers, output_shapes) ( onnx_io_output, @@ -355,18 +334,17 @@ class Gpt2BeamSearchTester(Gpt2Tester): return_numpy=False, include_copy_output_latency=True, ) + onnx_io_metric.add_latency(past_seq_len, avg_latency_ms / 1000.0) if test_data_saved < save_test_data: - onnx_io_runner.save_test_data( - session, onnx_io_output, save_test_data_dir, test_data_saved - ) + onnx_io_runner.save_test_data(session, onnx_io_output, save_test_data_dir, test_data_saved) test_data_saved += 1 onnx_io_runner.update(onnx_io_output, step, device) - done = done | (not onnx_runner.input_unfinished_sents.all()) - if torch.all(done): + if ((not onnx_runner.input_unfinished_sents.any()) + or (not torch_runner.input_unfinished_sents.any())): print("break at step: ", step) break @@ -379,15 +357,21 @@ class Gpt2BeamSearchTester(Gpt2Tester): onnx_io_metric.print() print("\tONNX") + if model_class == "GPT2LMHeadModel_BeamSearchStep": + results_onnx = onnx_runner.prev_step_results.view(batch_size * beam_size, -1) + results_onnx_io = onnx_io_runner.prev_step_results.view(batch_size * beam_size, -1) + else: + results_onnx = onnx_runner.input_ids.view(batch_size * beam_size, -1) + results_onnx_io = onnx_io_runner.input_ids.view(batch_size * beam_size, -1) Gpt2BeamSearchTester.pprint_results( - onnx_runner.prev_step_results.view(batch_size * beam_size, -1), + results_onnx, onnx_runner.prev_step_scores.view(batch_size * beam_size, -1), pad_token_id=eos_token_id, eos_token_id=eos_token_id, ) print("\tONNX with IO binding") Gpt2BeamSearchTester.pprint_results( - onnx_io_runner.prev_step_results.view(batch_size * beam_size, -1), + results_onnx_io, onnx_io_runner.prev_step_scores.view(batch_size * beam_size, -1), pad_token_id=eos_token_id, eos_token_id=eos_token_id, @@ -421,7 +405,7 @@ class Gpt2BeamSearchTester(Gpt2Tester): # remove EOS for k, t in enumerate(seq): if t == eos_token_id: - seq = seq[: k + 1] + seq = seq[:k + 1] break print("-" * 40) result = ",".join([str(token_id) for token_id in sample]) diff --git a/onnxruntime/python/tools/transformers/test/test_gpt2.py b/onnxruntime/python/tools/transformers/test/test_gpt2.py index 8f1258ed71..5b55ab227c 100644 --- a/onnxruntime/python/tools/transformers/test/test_gpt2.py +++ b/onnxruntime/python/tools/transformers/test/test_gpt2.py @@ -34,6 +34,37 @@ class TestGpt2(unittest.TestCase): def test_gpt2_int8(self): self.run_benchmark_gpt2('-m gpt2 --precision int8 -o -b 1 -s 128') + @pytest.mark.slow + def test_gpt2_beam_search_step_fp32(self): + self.run_benchmark_gpt2('-m gpt2 --model_class=GPT2LMHeadModel_BeamSearchStep --precision fp32 -v -b 1 -s 128') + + @pytest.mark.slow + def test_gpt2_beam_search_step_fp16(self): + if 'CUDAExecutionProvider' in onnxruntime.get_available_providers(): + self.run_benchmark_gpt2( + '-m gpt2 --model_class=GPT2LMHeadModel_BeamSearchStep --precision fp16 -o -b 1 -s 128 --use_gpu') + + @pytest.mark.slow + def test_gpt2_beam_search_step_int8(self): + self.run_benchmark_gpt2('-m gpt2 --model_class=GPT2LMHeadModel_BeamSearchStep --precision int8 -o -b 1 -s 128') + + @pytest.mark.slow + def test_gpt2_configurable_one_step_search_fp32(self): + self.run_benchmark_gpt2( + '-m gpt2 --model_class=GPT2LMHeadModel_ConfigurableOneStepSearch --precision fp32 -v -b 1 -s 128') + + @pytest.mark.slow + def test_gpt2_configurable_one_step_search_fp16(self): + if 'CUDAExecutionProvider' in onnxruntime.get_available_providers(): + self.run_benchmark_gpt2( + '-m gpt2 --model_class=GPT2LMHeadModel_ConfigurableOneStepSearch --precision fp16 -o -b 1 -s 128 --use_gpu' + ) + + @pytest.mark.slow + def test_gpt2_configurable_one_step_search_int8(self): + self.run_benchmark_gpt2( + '-m gpt2 --model_class=GPT2LMHeadModel_ConfigurableOneStepSearch --precision int8 -o -b 1 -s 128') + if __name__ == '__main__': import sys