Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
benchmark.py701 linesDownload Raw Back to llama
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 
codekingpro/portable-devtools · Team Ai