codekingpro/portable-devtools
114k
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License. See License.txt in the project root for
4# license information.
5# --------------------------------------------------------------------------
6import argparse
7import datetime
8import gc
9import itertools
10import logging
11import os
12import sys
13import time
14
15import numpy as np
16import onnx
17import psutil
18import torch
19from benchmark_helper import measure_memory, setup_logger
20from dist_settings import get_rank, get_size
21from llama_inputs import (
22 add_io_bindings_as_ortvalues,
23 get_merged_sample_with_past_kv_inputs,
24 get_msft_sample_inputs,
25 get_sample_inputs,
26 get_sample_with_past_kv_inputs,
27 verify_ort_inputs,
28)
29from optimum.onnxruntime import ORTModelForCausalLM
30from torch.profiler import ProfilerActivity, profile, record_function
31from tqdm import trange
32from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer
33
34import onnxruntime as ort
35
36logger = logging.getLogger(__name__)
37
38
39# For determining whether the ONNX model can do both prompt generation and token generation or only one of the two
40def get_ort_model_inputs_len(args, model):
41 if args.benchmark_type in {"hf-pt-eager", "hf-pt-compile"}:
42 return 0
43 if args.benchmark_type == "hf-ort":
44 try:
45 # New Optimum export (https://github.com/huggingface/optimum/blob/888332364c2e0091da1fc974737c7e277af168bf/optimum/onnxruntime/modeling_ort.py#L268)
46 return len(model.inputs_names)
47 except Exception:
48 # Old Optimum export (https://github.com/huggingface/optimum/blob/c5ad7f971cb0a494eac03dc0909f146725f999c5/optimum/onnxruntime/base.py#L54)
49 return len(model.decoder.input_names)
50 return len(model.get_inputs())
51
52
53def get_inputs(args: argparse.Namespace, ort_model_inputs_len: int):
54 init_inputs, iter_inputs = None, None
55
56 # For past_present_share_buffer:
57 # Set max_seq_len to 2048 for Microsoft LLaMA-2 model since that is the max value currently supported
58 # Set max_seq_len to config value for other models
59 max_seq_len = 2048 if args.benchmark_type == "ort-msft" else args.config.max_position_embeddings
60
61 if args.benchmark_type in {"hf-pt-eager", "hf-pt-compile"}:
62 init_inputs = get_sample_inputs(
63 args.config,
64 args.target_device,
65 args.batch_size,
66 args.sequence_length,
67 return_dict=True,
68 )
69 iter_inputs = get_sample_with_past_kv_inputs(
70 args.config,
71 args.target_device,
72 args.batch_size,
73 args.sequence_length,
74 use_fp16=args.use_fp16,
75 return_dict=True,
76 )
77
78 elif args.benchmark_type in {"hf-ort"}:
79 if ort_model_inputs_len == 3: # [input_ids, attention_mask, position_ids]
80 # Using split models in Optimum (e.g. created by Optimum export)
81 init_inputs = get_sample_inputs(
82 args.config,
83 args.target_device,
84 args.batch_size,
85 args.sequence_length,
86 return_dict=True,
87 )
88 iter_inputs = get_sample_with_past_kv_inputs(
89 args.config,
90 args.target_device,
91 args.batch_size,
92 args.sequence_length,
93 use_fp16=args.use_fp16,
94 return_dict=True,
95 )
96 else:
97 # Using merged model in Optimum (e.g. created by convert_to_onnx export)
98 init_inputs = get_merged_sample_with_past_kv_inputs(
99 args.config,
100 args.target_device,
101 args.batch_size,
102 seq_len=args.sequence_length,
103 past_seq_len=0,
104 max_seq_len=max_seq_len,
105 use_fp16=args.use_fp16,
106 use_buffer_share=args.use_buffer_share,
107 engine="pt",
108 return_dict=True,
109 )
110 iter_inputs = get_merged_sample_with_past_kv_inputs(
111 args.config,
112 args.target_device,
113 args.batch_size,
114 seq_len=1,
115 past_seq_len=args.sequence_length,
116 max_seq_len=max_seq_len,
117 use_fp16=args.use_fp16,
118 use_buffer_share=args.use_buffer_share,
119 engine="pt",
120 return_dict=True,
121 )
122
123 elif args.benchmark_type == "ort-convert-to-onnx":
124 # Microsoft export from convert_to_onnx
125 init_inputs = get_merged_sample_with_past_kv_inputs(
126 args.config,
127 args.target_device,
128 args.batch_size,
129 seq_len=args.sequence_length,
130 past_seq_len=0,
131 max_seq_len=max_seq_len,
132 use_fp16=args.use_fp16,
133 use_buffer_share=args.use_buffer_share,
134 engine="ort",
135 return_dict=True,
136 world_size=args.world_size,
137 )
138 iter_inputs = get_merged_sample_with_past_kv_inputs(
139 args.config,
140 args.target_device,
141 args.batch_size,
142 seq_len=1,
143 past_seq_len=args.sequence_length,
144 max_seq_len=max_seq_len,
145 use_fp16=args.use_fp16,
146 use_buffer_share=args.use_buffer_share,
147 engine="ort",
148 return_dict=True,
149 world_size=args.world_size,
150 )
151
152 elif args.benchmark_type == "ort-msft":
153 # Microsoft export from https://github.com/microsoft/Llama-2-Onnx
154 split_kv = ort_model_inputs_len > 5 # original inputs: [x, attn_mask, k_cache, v_cache, pos]
155
156 init_inputs = get_msft_sample_inputs(
157 args.config,
158 args.batch_size,
159 past_seq_len=0,
160 seq_len=args.sequence_length,
161 max_seq_len=max_seq_len,
162 use_fp16=args.use_fp16,
163 use_buffer_share=args.use_buffer_share,
164 split_kv=split_kv,
165 )
166 iter_inputs = get_msft_sample_inputs(
167 args.config,
168 args.batch_size,
169 past_seq_len=args.sequence_length,
170 seq_len=1,
171 max_seq_len=max_seq_len,
172 use_fp16=args.use_fp16,
173 use_buffer_share=args.use_buffer_share,
174 split_kv=split_kv,
175 )
176
177 else:
178 raise Exception("Unable to auto-detect inputs for provided model")
179
180 return init_inputs, iter_inputs
181
182
183def get_model(args: argparse.Namespace):
184 model, sess_options = None, None
185 start_time, end_time = None, None
186
187 # There are multiple sources that the model could come from:
188 # 1) Benchmark LLaMA-2 from unofficial source on Hugging Face
189 # 2) Benchmark LLaMA-2 from official source on Hugging Face, which requires an authentication token
190 # 3) Benchmark LLaMA-2 from local download of model
191 # 4) Benchmark LLaMA-2 from Microsoft (already optimized, available at https://github.com/microsoft/Llama-2-Onnx)
192 # 5) Benchmark LLaMA-2 from convert_to_onnx
193
194 if args.benchmark_type in {"hf-pt-eager", "hf-pt-compile"}:
195 source = args.hf_pt_dir_path if args.hf_pt_dir_path else args.model_name
196 start_time = time.time()
197 model = AutoModelForCausalLM.from_pretrained(
198 source,
199 torch_dtype=torch.float16 if args.use_fp16 else torch.float32,
200 use_auth_token=args.auth,
201 trust_remote_code=args.auth,
202 use_cache=True,
203 cache_dir=args.cache_dir,
204 ).to(args.target_device)
205 end_time = time.time()
206
207 if args.benchmark_type == "hf-pt-compile":
208 model = torch.compile(model)
209
210 elif args.benchmark_type in {"hf-ort", "ort-msft", "ort-convert-to-onnx"}:
211 sess_options = ort.SessionOptions()
212 sess_options.enable_profiling = args.profile
213 if args.verbose:
214 sess_options.log_verbosity_level = 1
215 sess_options.log_severity_level = 1
216
217 else:
218 raise Exception(f"Cannot recognize {args.benchmark_type}")
219
220 if args.benchmark_type == "hf-ort":
221 # Optimum export or convert_to_onnx.py export
222 provider = args.execution_provider[0] if type(args.execution_provider) is tuple else args.execution_provider
223 provider_options = args.execution_provider[1] if type(args.execution_provider) is tuple else None
224
225 decoder_file_name = None
226 decoder_with_past_file_name = None
227 for filename in os.listdir(args.hf_ort_dir_path):
228 if ".onnx" not in filename or ".onnx_data" in filename or ".onnx.data" in filename:
229 continue
230 if "decoder_model" in filename or filename == "model.onnx":
231 decoder_file_name = filename
232 if "decoder_with_past_model" in filename:
233 decoder_with_past_file_name = filename
234 if "decoder_merged_model" in filename:
235 decoder_file_name = filename
236 decoder_with_past_file_name = filename
237
238 start_time = time.time()
239 model = ORTModelForCausalLM.from_pretrained(
240 args.hf_ort_dir_path,
241 decoder_file_name=decoder_file_name,
242 decoder_with_past_file_name=decoder_with_past_file_name,
243 use_auth_token=args.auth,
244 trust_remote_code=args.auth,
245 use_io_binding=True, # Large perf gain even for cpu due to avoiding output copy.
246 use_merged=(True if decoder_file_name == "model.onnx" else None),
247 provider=provider,
248 provider_options=provider_options,
249 session_options=sess_options,
250 )
251 end_time = time.time()
252
253 if args.benchmark_type in {"ort-msft", "ort-convert-to-onnx"}:
254 # Ex: Microsoft export from https://github.com/microsoft/Llama-2-Onnx
255 logger.info(f"Loading model from {args.ort_model_path.format(args.rank)}")
256 start_time = time.time()
257 model = ort.InferenceSession(
258 args.ort_model_path.format(args.rank),
259 sess_options,
260 providers=[args.execution_provider],
261 )
262 end_time = time.time()
263
264 logger.info(f"Loaded model in {end_time - start_time} s")
265 return model
266
267
268def time_fn(args, fn, inputs):
269 # Warm up
270 warmup_range = (
271 range(args.warmup_runs)
272 if args.benchmark_type in {"ort-msft", "ort-convert-to-onnx"}
273 else trange(args.warmup_runs, file=sys.stdout, desc="Warm up")
274 )
275
276 if args.verbose:
277 outputs = fn(inputs)
278 logger.info(outputs)
279
280 input_sync = lambda *kwargs: ( # noqa: E731
281 args.io_binding.synchronize_inputs()
282 if args.device != "cpu" and args.benchmark_type in {"ort-msft", "ort-convert-to-onnx"} # ORT synchronize
283 else lambda *kwargs: (
284 torch.cuda.synchronize()
285 if args.device != "cpu" and torch.cuda.is_available() # PyTorch synchronize
286 else lambda *kwargs: None
287 )
288 ) # no-op function
289
290 output_sync = lambda *kwargs: ( # noqa: E731
291 args.io_binding.synchronize_outputs()
292 if args.device != "cpu" and args.benchmark_type in {"ort-msft", "ort-convert-to-onnx"} # ORT synchronize
293 else lambda *kwargs: (
294 torch.cuda.synchronize()
295 if args.device != "cpu" and torch.cuda.is_available() # PyTorch synchronize
296 else lambda *kwargs: None
297 )
298 ) # no-op function
299
300 for _ in warmup_range:
301 input_sync()
302 fn(inputs)
303 output_sync()
304
305 # Benchmark
306 total_time = 0
307 bench_range = (
308 range(args.num_runs)
309 if args.benchmark_type in {"ort-msft", "ort-convert-to-onnx"}
310 else trange(args.num_runs, file=sys.stdout, desc="Benchmark")
311 )
312 for _ in bench_range:
313 input_sync()
314 start_time = time.time()
315
316 fn(inputs)
317
318 output_sync()
319 end_time = time.time()
320
321 total_time += end_time - start_time
322
323 # Newline print after trange in order to print metrics on new lines without progress bar on same line
324 if args.benchmark_type not in {"ort-msft", "ort-convert-to-onnx"}:
325 logger.info("")
326
327 latency = total_time / args.num_runs
328 throughput = args.batch_size / latency
329
330 if args.rank == 0:
331 logger.info(f"Batch Size: {args.batch_size}")
332 logger.info(f"Sequence Length: {args.sequence_length}")
333 logger.info(f"Latency: {latency} s")
334 logger.info(f"Throughput: {throughput} tps")
335 return
336
337
338def profile_fn(args, fn, inputs, inputs_type):
339 # Filename prefix format:
340 # "b<batch-size>_s<sequence-length>_<benchmark-type>-<precision>-<device>_<inference-step>_<inputs-type>_<current-time>"
341 prefix = f"b{args.batch_size}_s{args.sequence_length}_{args.benchmark_type.lower()}-{args.precision}-{args.device}_{fn.__name__.replace('_', '-')}_{inputs_type}_{datetime.datetime.now():%Y-%m-%d_%H:%M:%S}"
342 filename = None
343
344 if args.benchmark_type in {"hf-pt-eager", "hf-pt-compile"}:
345 # Profile PyTorch kernels
346 with profile( # noqa: SIM117
347 activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA], record_shapes=True, profile_memory=True
348 ) as prof:
349 with record_function("model_inference"):
350 fn(inputs)
351 prof_data = prof.key_averages(group_by_stack_n=5).table(sort_by=args.pt_filter_by, row_limit=args.pt_num_rows)
352
353 filename = os.path.join(args.log_folder, f"{prefix}.log")
354 with open(filename, "w") as f:
355 f.write(prof_data)
356
357 else:
358 # Profile ORT kernels
359 fn(inputs)
360
361 # Set new log name for ORT profile log generated
362 filename = f"{prefix}.json"
363
364 return filename
365
366
367def measure_fn(args, fn, inputs):
368 # Measure CPU usage
369 pid = os.getpid()
370 process = psutil.Process(pid)
371 process.cpu_percent(interval=0.1)
372
373 fn(inputs)
374 if args.rank == 0:
375 logger.info(f"CPU usage: {process.cpu_percent(interval=None) / psutil.cpu_count(logical=False)}%")
376
377 # Measure memory usage
378 gc.collect()
379 torch.cuda.empty_cache()
380 measure_memory(is_gpu=(args.device != "cpu"), func=lambda: fn(inputs))
381
382 # Flush output so memory usage is printed
383 sys.stdout.flush()
384
385
386def run_hf_inference(args, init_inputs, iter_inputs, model):
387 # Inference steps to measure
388 def get_logits(inputs):
389 # Inference pass without decoding
390 outputs = model(**inputs)
391 return outputs
392
393 # Examples of other inference steps that can be measured:
394 # To use, uncomment the function and assign it to `generate_fn`
395
396 # def get_pred_ids(inputs):
397 # # Inference pass with predicted token ids generation
398 # predicted_ids = model.generate(**inputs)
399 # return predicted_ids
400
401 # def gen_and_dec(inputs):
402 # # Inference pass with generation and decoding
403 # predicted_ids = get_pred_ids(inputs)
404 # transcription = []
405 # for bs in range(args.batch_size):
406 # for rs in range(args.num_return_sequences):
407 # transcription.append(
408 # args.tokenizer.batch_decode(
409 # predicted_ids[bs * args.num_return_sequences + rs], skip_special_tokens=True
410 # )[0]
411 # )
412 # return transcription
413
414 generate_fn = get_logits
415
416 if args.benchmark_type == "hf-pt-compile":
417 # Run forward pass once with each set of inputs to process through Dynamo
418 generate_fn(init_inputs)
419 generate_fn(iter_inputs)
420
421 if args.profile:
422 new_logname = profile_fn(args, generate_fn, init_inputs, "prompt")
423 if args.benchmark_type == "hf-ort":
424 # Turn profiling off to stop appending to log
425 old_logname = model.decoder.session.end_profiling()
426 logger.warning(f"Renaming {old_logname} to {new_logname}")
427 os.rename(old_logname, os.path.join(args.log_folder, new_logname))
428
429 new_logname = profile_fn(args, generate_fn, iter_inputs, "token")
430 if args.benchmark_type == "hf-ort":
431 # Turn profiling off to stop appending to log
432 old_logname = model.decoder_with_past.session.end_profiling()
433 logger.warning(f"Renaming {old_logname} to {new_logname}")
434 os.rename(old_logname, os.path.join(args.log_folder, new_logname))
435
436 return
437
438 # PyTorch evaluations
439 logger.info("\nEvaluating `model(inputs)` step to get past_key_values")
440 time_fn(args, generate_fn, init_inputs)
441 measure_fn(args, generate_fn, init_inputs)
442
443 logger.info("\nEvaluating `model(inputs)` step with past_key_values")
444 time_fn(args, generate_fn, iter_inputs)
445 measure_fn(args, generate_fn, iter_inputs)
446
447
448def run_ort_inference(args, init_inputs, iter_inputs, model):
449 def prepare_ort_inputs(inputs, kv_cache_ortvalues):
450 # Verify model inputs
451 inputs = verify_ort_inputs(model, inputs)
452
453 # Add IO bindings for non-CPU execution providers
454 if args.device != "cpu":
455 io_binding, kv_cache_ortvalues = add_io_bindings_as_ortvalues(
456 model, inputs, args.device, int(args.rank), args.use_buffer_share, kv_cache_ortvalues
457 )
458 setattr(args, "io_binding", io_binding) # noqa: B010
459 return io_binding, kv_cache_ortvalues
460
461 return inputs, kv_cache_ortvalues
462
463 def with_io_binding(io_binding):
464 # Inference pass with IO binding
465 model.run_with_iobinding(io_binding)
466
467 def without_io_binding(inputs):
468 # Inference pass without IO binding
469 outputs = model.run(None, inputs)
470 return outputs
471
472 generate_fn = with_io_binding if args.device != "cpu" else without_io_binding
473 kv_cache_ortvalues = {}
474
475 if args.profile:
476 ort_init_inputs, kv_cache_ortvalues = prepare_ort_inputs(init_inputs, kv_cache_ortvalues)
477 new_logname = profile_fn(args, generate_fn, ort_init_inputs, "prompt")
478
479 # Turn profiling off to stop appending to log file
480 old_logname = model.end_profiling()
481 logger.warning(f"Renaming {old_logname} to {new_logname}")
482 os.rename(old_logname, os.path.join(args.log_folder, new_logname))
483
484 # Re-initialize model for new log file instead of appending to old log file
485 model = get_model(args)
486 ort_iter_inputs, kv_cache_ortvalues = prepare_ort_inputs(iter_inputs, kv_cache_ortvalues)
487 new_logname = profile_fn(args, generate_fn, ort_iter_inputs, "token")
488
489 # Turn profiling off to stop appending to log
490 old_logname = model.end_profiling()
491 logger.warning(f"Renaming {old_logname} to {new_logname}")
492 os.rename(old_logname, os.path.join(args.log_folder, new_logname))
493 return
494
495 # ORT evaluations
496 logger.info("\nEvaluating `model(inputs)` step to get past_key_values")
497 ort_init_inputs, kv_cache_ortvalues = prepare_ort_inputs(init_inputs, kv_cache_ortvalues)
498 time_fn(args, generate_fn, ort_init_inputs)
499 measure_fn(args, generate_fn, ort_init_inputs)
500
501 logger.info("\nEvaluating `model(inputs)` step with past_key_values")
502 ort_iter_inputs, kv_cache_ortvalues = prepare_ort_inputs(iter_inputs, kv_cache_ortvalues)
503 time_fn(args, generate_fn, ort_iter_inputs)
504 measure_fn(args, generate_fn, ort_iter_inputs)
505
506
507def run_inference(args, init_inputs, iter_inputs, model):
508 if args.benchmark_type in {"hf-pt-eager", "hf-pt-compile", "hf-ort"}:
509 run_hf_inference(args, init_inputs, iter_inputs, model)
510 elif args.benchmark_type in {"ort-msft", "ort-convert-to-onnx"}:
511 run_ort_inference(args, init_inputs, iter_inputs, model)
512 else:
513 raise Exception(f"Cannot recognize {args.benchmark_type}")
514
515
516def get_args(rank=0):
517 parser = argparse.ArgumentParser()
518 parser.add_argument(
519 "-bt",
520 "--benchmark-type",
521 type=str,
522 required=True,
523 choices=[
524 "hf-pt-eager",
525 "hf-pt-compile",
526 "hf-ort",
527 "ort-msft",
528 "ort-convert-to-onnx",
529 ],
530 )
531 parser.add_argument(
532 "-m",
533 "--model-name",
534 type=str,
535 required=True,
536 help="Hugging Face name of model (e.g. 'meta-llama/Llama-2-7b-hf')",
537 )
538 parser.add_argument(
539 "-a", "--auth", default=False, action="store_true", help="Use Hugging Face authentication token to access model"
540 )
541
542 # Args for choosing the model
543 parser.add_argument(
544 "-p",
545 "--precision",
546 required=True,
547 type=str,
548 default="fp32",
549 choices=["int4", "int8", "fp16", "fp32"],
550 help="Precision for model. For ONNX models, the model's precision should be set before running this script.",
551 )
552 parser.add_argument(
553 "--hf-pt-dir-path",
554 type=str,
555 default="",
556 help="Path to directory containing all PyTorch files (e.g. tokenizer, PyTorch model)",
557 )
558 parser.add_argument(
559 "--hf-ort-dir-path",
560 type=str,
561 default="",
562 help="Path to directory containing all ONNX files (e.g. tokenizer, decoder_merged, decoder, decoder_with_past)",
563 )
564 parser.add_argument(
565 "--ort-model-path",
566 type=str,
567 default="",
568 help="Path to ONNX model",
569 )
570
571 # Args for running and evaluating the model
572 parser.add_argument(
573 "-b",
574 "--batch-sizes",
575 default="1 2",
576 )
577 parser.add_argument(
578 "-s",
579 "--sequence-lengths",
580 default="32 64 128 256 512",
581 )
582 parser.add_argument(
583 "-d",
584 "--device",
585 type=str,
586 default="cuda" if torch.cuda.is_available() else "cpu",
587 choices=["cpu", "cuda"],
588 )
589 parser.add_argument("-id", "--device-id", type=int, default=0)
590 parser.add_argument("-w", "--warmup-runs", type=int, default=5)
591 parser.add_argument("-n", "--num-runs", type=int, default=10)
592 parser.add_argument("--seed", type=int, default=2)
593
594 # Args for decoding logic
595 parser.add_argument("--max-length", type=int, default=32)
596 parser.add_argument("--num-return-sequences", type=int, default=1)
597
598 # Args for accessing detailed info
599 parser.add_argument("--profile", default=False, action="store_true")
600 parser.add_argument(
601 "--pt-filter-by", type=str, default="self_cpu_time_total", help="What to filter PyTorch profiler by"
602 )
603 parser.add_argument("--pt-num-rows", type=int, default=1000, help="Number of rows for PyTorch profiler to display")
604 parser.add_argument("--verbose", default=False, action="store_true")
605 parser.add_argument("--log-folder", type=str, default=os.path.join("."), help="Folder to cache log files")
606 parser.add_argument(
607 "--cache-dir",
608 type=str,
609 required=True,
610 default="./model_cache",
611 help="Cache dir where Hugging Face files are stored",
612 )
613
614 args = parser.parse_args()
615
616 # Set seed properties
617 np.random.seed(args.seed)
618 torch.manual_seed(args.seed)
619
620 # Set runtime properties
621 if "ort" in args.benchmark_type:
622 setattr(args, "execution_provider", f"{args.device.upper()}ExecutionProvider") # noqa: B010
623 if args.execution_provider == "CUDAExecutionProvider":
624 args.execution_provider = (args.execution_provider, {"device_id": rank})
625
626 # Check that paths have been specified for any benchmarking with ORT
627 if args.benchmark_type == "hf-ort":
628 assert args.hf_ort_dir_path, "Please specify a path to `--hf-ort-dir-path`"
629 if args.benchmark_type in {"ort-msft", "ort-convert-to-onnx"}:
630 assert args.ort_model_path, "Please specify a path to `--ort-model-path`"
631
632 args.batch_sizes = args.batch_sizes.split(" ")
633 args.sequence_lengths = args.sequence_lengths.split(" ")
634
635 # Use FP32 precision for FP32, INT8, INT4 CPU models, use FP16 precision for FP16 and INT4 GPU models
636 args.precision = (
637 "fp32" if args.precision in {"int8", "fp32"} or (args.precision == "int4" and args.device == "cpu") else "fp16"
638 )
639
640 # Check that only one (batch_size, sequence_length) combination is set for profiling
641 if args.profile:
642 assert len(args.batch_sizes) == 1 and len(args.sequence_lengths) == 1, (
643 "Please provide only one (batch_size, sequence_length) combination for profiling"
644 )
645
646 return args
647
648
649def main():
650 rank = get_rank()
651 world_size = get_size()
652
653 args = get_args(rank)
654 setup_logger(args.verbose)
655 logger.info(args.__dict__)
656 torch.backends.cudnn.benchmark = True
657
658 args.rank = rank
659 args.world_size = world_size
660 tokenizer = AutoTokenizer.from_pretrained(
661 args.model_name, cache_dir=args.cache_dir, use_auth_token=args.auth, trust_remote_code=args.auth
662 )
663 config = AutoConfig.from_pretrained(
664 args.model_name, cache_dir=args.cache_dir, use_auth_token=args.auth, trust_remote_code=args.auth
665 )
666 target_device = f"cuda:{args.rank}" if args.device != "cpu" else args.device
667 use_fp16 = args.precision == "fp16"
668
669 setattr(args, "tokenizer", tokenizer) # noqa: B010
670 setattr(args, "config", config) # noqa: B010
671 setattr(args, "target_device", target_device) # noqa: B010
672 setattr(args, "use_fp16", use_fp16) # noqa: B010
673
674 # Get model and model info
675 model = get_model(args)
676 ort_model_inputs_len = get_ort_model_inputs_len(args, model)
677
678 # Check if past_present_share_buffer can be enabled (only for FP16 models with GQA)
679 if args.benchmark_type in {"ort-convert-to-onnx", "ort-msft"}:
680 onnx_model = onnx.load_model(args.ort_model_path.format(args.rank), load_external_data=False)
681 gqa_nodes = list(filter(lambda node: node.op_type == "GroupQueryAttention", onnx_model.graph.node))
682
683 use_buffer_share = use_fp16 and len(gqa_nodes) > 0 and args.device != "cpu"
684 setattr(args, "use_buffer_share", use_buffer_share) # noqa: B010
685 else:
686 setattr(args, "use_buffer_share", False) # noqa: B010
687
688 # Measure prompt cost (init_inputs) and generated token cost (iter_inputs)
689 for batch_size, sequence_length in itertools.product(args.batch_sizes, args.sequence_lengths):
690 if args.rank == 0:
691 logger.info(f"\nBatch size = {batch_size} and sequence length = {sequence_length}...")
692 setattr(args, "batch_size", int(batch_size)) # noqa: B010
693 setattr(args, "sequence_length", int(sequence_length)) # noqa: B010
694
695 init_inputs, iter_inputs = get_inputs(args, ort_model_inputs_len)
696 run_inference(args, init_inputs, iter_inputs, model)
697
698
699if __name__ == "__main__":
700 main()
701 