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 <xiaoyu@xiaoyu-VM.z4vh1dzj5eoevgybsksdpz2izh.jx.internal.cloudapp.net>
This commit is contained in:
Xiaoyu Liu 2021-04-29 13:19:56 -07:00 committed by GitHub
parent 6358e96b63
commit 994c2ed420
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
5 changed files with 679 additions and 386 deletions

View file

@ -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)

View file

@ -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)

View file

@ -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

View file

@ -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])

View file

@ -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