Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
benchmark_e2e.py609 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# --------------------------------------------------------------------------
6
7# This is an end-to-end benchmarking script for the Hugging Face LLaMA-2 model.
8#
9# Prerequisites:
10# 1) Install `huggingface-cli`:
11#
12# $ pip install huggingface_hub
13#
14# 2) Authenticate with Hugging Face's CLI:
15#
16# $ huggingface-cli login
17#
18# 3) Accept Meta's license in Hugging Face to access the models at https://huggingface.co/meta-llama/
19#
20# 4) Install the latest ONNX Runtime version
21#
22# $ pip install onnxruntime-gpu
23#
24# 5) Install flash attention v2
25#
26# $ pip install flash-attn --no-build-isolation
27#
28# 6) Install bitsandbytes
29#
30# $ pip install bitsandbytes
31
32from __future__ import annotations
33
34import argparse
35import datetime
36import gc
37import itertools
38import json
39import logging
40import os
41import textwrap
42import time
43
44import numpy as np
45import pandas as pd
46import torch
47from benchmark_helper import setup_logger
48from llama_inputs import add_io_bindings_as_tensors, get_initial_inputs_and_outputs
49from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
50
51import onnxruntime as ort
52
53logger = logging.getLogger(__name__)
54
55
56def get_model(args: argparse.Namespace):
57    if args.benchmark_type in {"pt-eager", "pt-compile"}:
58        model = None
59        if args.onnx_precision == "int4" and args.device == "cuda":
60            bnb_config = BitsAndBytesConfig(
61                load_in_4bit=True,
62                bnb_4bit_use_double_quant=True,
63                bnb_4bit_quant_type="nf4",
64                bnb_4bit_compute_dtype=torch.float16,
65            )
66
67            model = AutoModelForCausalLM.from_pretrained(
68                args.hf_dir_path if args.hf_dir_path != "" else args.model_name,
69                cache_dir=args.cache_dir,
70                torch_dtype=args.torch_dtype,
71                use_auth_token=args.auth,
72                trust_remote_code=args.trust,
73                use_cache=True,
74                attn_implementation="flash_attention_2",
75                quantization_config=bnb_config,
76                max_memory={args.device_id: "80GB"},
77            )
78        else:
79            try:
80                model = AutoModelForCausalLM.from_pretrained(
81                    args.hf_dir_path if args.hf_dir_path != "" else args.model_name,
82                    cache_dir=args.cache_dir,
83                    torch_dtype=args.torch_dtype,
84                    use_auth_token=args.auth,
85                    trust_remote_code=args.trust,
86                    use_cache=True,
87                    attn_implementation=("flash_attention_2" if args.device == "cuda" else "sdpa"),
88                ).to(args.target_device)
89            except Exception as e:
90                # When flash_attention or sdpa doesn't support a model, it throws an exception.
91                # Rather than stopping a process, run as eager mode.
92                print("Try to load a model using eager mode: ", e)
93                model = AutoModelForCausalLM.from_pretrained(
94                    args.hf_dir_path if args.hf_dir_path != "" else args.model_name,
95                    cache_dir=args.cache_dir,
96                    torch_dtype=args.torch_dtype,
97                    use_auth_token=args.auth,
98                    trust_remote_code=args.trust,
99                    use_cache=True,
100                    attn_implementation="eager",
101                ).to(args.target_device)
102
103        model.eval()
104
105        if args.benchmark_type == "pt-compile":
106            model = torch.compile(model)
107
108    else:
109        sess_options = ort.SessionOptions()
110        ep = (
111            ("CUDAExecutionProvider", {"device_id": args.device_id})
112            if args.device == "cuda"
113            else "CPUExecutionProvider"
114        )
115        model = ort.InferenceSession(args.onnx_model_path, sess_options=sess_options, providers=[ep])
116
117    return model
118
119
120def run_inference(args, model, runs, inputs, outputs):
121    if args.benchmark_type == "pt-compile":
122        with torch.no_grad():
123            outputs = model(**inputs)
124
125    # Synchronize inputs
126    io_binding = None
127    if args.benchmark_type in {"pt-eager", "pt-compile"}:
128        if args.device != "cpu":
129            torch.cuda.synchronize(args.target_device)
130    else:
131        io_binding = add_io_bindings_as_tensors(model, inputs, outputs, args.use_fp16, args.use_buffer_share)
132        io_binding.synchronize_inputs()
133
134    # Run inference
135    start = time.perf_counter()
136    for _ in range(runs):
137        if args.benchmark_type in {"pt-eager", "pt-compile"}:
138            with torch.no_grad():
139                outputs = model(**inputs)
140                if args.device != "cpu":
141                    torch.cuda.synchronize(args.target_device)
142        else:
143            model.run_with_iobinding(io_binding)
144            io_binding.synchronize_outputs()
145
146    end = time.perf_counter()
147    avg = (end - start) / runs
148    return avg, outputs
149
150
151def prepare_model_for_inference(args, model, config, tokenizer, prompt_length, prompt):
152    clear_cache()
153    inputs, outputs = get_initial_inputs_and_outputs(
154        config, tokenizer, prompt_length, prompt, args.target_device, args.use_fp16, args.use_buffer_share, args.engine
155    )
156    _, outputs = run_inference(args, model, args.warmup_runs, inputs, outputs)
157    return inputs, outputs
158
159
160def clear_cache():
161    gc.collect()
162    torch.cuda.empty_cache()
163
164
165def save_results(results, filename, gen_length):
166    df = pd.DataFrame(
167        results,
168        columns=[
169            "Batch Size",
170            "Prompt Length",
171            "Prompt Processing Latency (ms)",
172            "Prompt Processing Throughput (tps)",
173            "Sampling Latency (ms)",
174            "Sampling Throughput (tps)",
175            "First Token Generated Latency (ms)",
176            "First Token Generated Throughput (tps)",
177            f"Average Latency of First {gen_length // 2} Tokens Generated (ms)",
178            f"Average Throughput of First {gen_length // 2} Tokens Generated (tps)",
179            f"Average Latency of First {gen_length} Tokens Generated (ms)",
180            f"Average Throughput of First {gen_length} Tokens Generated (tps)",
181            "Wall-Clock Latency (s)",
182            "Wall-Clock Throughput (tps)",
183        ],
184    )
185
186    df.to_csv(filename, index=False)
187    logger.info(f"Results saved in {filename}!")
188
189
190def get_args():
191    parser = argparse.ArgumentParser()
192
193    parser.add_argument(
194        "-bt",
195        "--benchmark-type",
196        type=str,
197        required=True,
198        choices=["pt-eager", "pt-compile", "ort"],
199    )
200
201    parser.add_argument(
202        "-m",
203        "--model-name",
204        type=str,
205        required=False,
206        help="Hugging Face name of model (e.g. 'meta-llama/Llama-2-7b-hf')",
207    )
208
209    parser.add_argument(
210        "-a",
211        "--auth",
212        default=False,
213        action="store_true",
214        help="Use Hugging Face authentication token to access model",
215    )
216
217    parser.add_argument(
218        "-t",
219        "--trust",
220        default=False,
221        action="store_true",
222        help="Whether or not to allow for custom models defined on the Hugging Face Hub in their own modeling files",
223    )
224
225    parser.add_argument(
226        "-c",
227        "--cache-dir",
228        type=str,
229        default=os.path.join(".", "model_cache"),
230        help="Path to directory containing all Hugging Face files (e.g. config, tokenizer, PyTorch model). Use when loading model as `AutoModel.from_pretrained(model_name, cache_dir=cache_dir)`.",
231    )
232
233    parser.add_argument(
234        "--hf-dir-path",
235        type=str,
236        default="",
237        help="Path to directory containing all Hugging Face files (e.g. config, tokenizer, PyTorch model). Use when loading model as `AutoModel.from_pretrained(folder_path)`.",
238    )
239
240    parser.add_argument(
241        "-o",
242        "--onnx-model-path",
243        required=False,
244        help="Path to ONNX model",
245    )
246
247    parser.add_argument(
248        "-f",
249        "--prompts-file",
250        required=True,
251        default=os.path.join(".", "models", "llama", "prompts.json"),
252        help="JSON file containing entries in the format 'prompt length: prompt' where prompt length = tokenized length of prompt",
253    )
254
255    parser.add_argument(
256        "--use_buffer_share",
257        default=False,
258        action="store_true",
259        help="Use when GroupQueryAttention (GQA) is in ONNX model",
260    )
261
262    (
263        parser.add_argument(
264            "--anomaly-filtering",
265            default=False,
266            action="store_true",
267            help="Use this flag to filter anomaly accelerator times for tokens generated. \
268              This may give more accurate latency and throughput metrics for tokens generated. \
269              Wall-clock metrics are still reported with anomaly times though.",
270        ),
271    )
272
273    parser.add_argument(
274        "-b",
275        "--batch-sizes",
276        default="1 2",
277    )
278
279    parser.add_argument(
280        "-s",
281        "--prompt-lengths",
282        default="16 64 256 1024",
283    )
284
285    parser.add_argument(
286        "-p",
287        "--precision",
288        required=True,
289        type=str,
290        default="fp32",
291        choices=["int4", "int8", "fp16", "fp32"],
292        help="Precision for model. For ONNX models, the model's precision should be set before running this script.",
293    )
294
295    parser.add_argument(
296        "-g",
297        "--generation-length",
298        type=int,
299        default=256,
300        help="Number of new tokens to generate",
301    )
302
303    parser.add_argument(
304        "-d",
305        "--device",
306        type=str,
307        default="cuda" if torch.cuda.is_available() else "cpu",
308        choices=["cpu", "cuda"],
309    )
310
311    parser.add_argument("-id", "--device-id", type=int, default=0)
312    parser.add_argument("-w", "--warmup-runs", type=int, default=5)
313    parser.add_argument("-n", "--num-runs", type=int, default=100)
314    parser.add_argument("--seed", type=int, default=2)
315
316    args = parser.parse_args()
317
318    # Set seed properties
319    np.random.seed(args.seed)
320    torch.manual_seed(args.seed)
321
322    # Set runtime properties
323    if "ort" in args.benchmark_type:
324        setattr(args, "execution_provider", f"{args.device.upper()}ExecutionProvider")  # noqa: B010
325        if args.execution_provider == "CUDAExecutionProvider":
326            args.execution_provider = (args.execution_provider, {"device_id": args.device_id})
327
328    # Check that paths have been specified for any benchmarking with ORT
329    if args.benchmark_type == "ort":
330        assert args.onnx_model_path, "Please specify a path to `--onnx-model-path`"
331
332    args.batch_sizes = args.batch_sizes.split(" ")
333    args.prompt_lengths = args.prompt_lengths.split(" ")
334
335    # Use FP32 precision for FP32, INT8, INT4 CPU models, use FP16 precision for FP16 and INT4 GPU models
336    setattr(args, "onnx_precision", args.precision)  # noqa: B010
337    args.precision = (
338        "fp32" if args.precision in {"int8", "fp32"} or (args.precision == "int4" and args.device == "cpu") else "fp16"
339    )
340
341    target_device = f"cuda:{args.device_id}" if args.device != "cpu" else args.device
342    torch_dtype = torch.float16 if args.precision == "fp16" else torch.float32
343    engine = "ort" if args.benchmark_type == "ort" else "pt"
344    setattr(args, "target_device", target_device)  # noqa: B010
345    setattr(args, "torch_dtype", torch_dtype)  # noqa: B010
346    setattr(args, "engine", engine)  # noqa: B010
347    setattr(args, "use_fp16", args.precision == "fp16")  # noqa: B010
348
349    args.use_buffer_share = args.use_buffer_share and engine == "ort"
350
351    return args
352
353
354def main():
355    args = get_args()
356    setup_logger(False)
357    logger.info(args.__dict__)
358
359    # Get prompts and prompt sizes
360    size_to_prompt = None
361    with open(args.prompts_file) as f:
362        size_to_prompt = json.load(f, object_hook=lambda d: {int(k): v for k, v in d.items()})
363
364    # Get config, tokenizer, and model
365    config = AutoConfig.from_pretrained(
366        args.hf_dir_path if args.hf_dir_path != "" else args.model_name,
367        cache_dir=args.cache_dir,
368        use_auth_token=args.auth,
369        trust_remote_code=args.trust,
370    )
371    tokenizer = AutoTokenizer.from_pretrained(
372        args.hf_dir_path if args.hf_dir_path != "" else args.model_name,
373        cache_dir=args.cache_dir,
374        use_auth_token=args.auth,
375        trust_remote_code=args.trust,
376    )
377    model = get_model(args)
378
379    all_csv_metrics = []
380    for batch_size, prompt_length in itertools.product(args.batch_sizes, args.prompt_lengths):
381        batch_size, prompt_length = int(batch_size), int(prompt_length)  # noqa: PLW2901
382        logger.info(f"Running batch size = {batch_size}, prompt length = {prompt_length}")
383        clear_cache()
384        max_length = prompt_length + args.generation_length
385
386        if prompt_length not in size_to_prompt:
387            raise NotImplementedError(
388                textwrap.dedent(
389                    f"""
390                                A prompt of size {prompt_length} was not found in '{args.prompts_file}'. There are a couple of solutions to fix this.
391                                1) You can change one of the keys in '{args.prompts_file}' to be {prompt_length}.
392                                    If {prompt_length} < actual prompt's length, the benchmark E2E tool will repeat the first word in the prompt until {prompt_length} = actual prompt's length.
393                                    If {prompt_length} > actual prompt's length, the benchmark E2E tool will automatically trim the actual prompt's length so that {prompt_length} = actual prompt's length.
394                                2) You can add a new key-value entry in '{args.prompts_file}' of the form '{prompt_length}': 'your prompt goes here'.
395                """
396                )
397            )
398        prompt = [size_to_prompt[prompt_length]] * batch_size
399        csv_metrics = [batch_size, prompt_length]
400
401        try:
402            # Measure prompt processing
403            logger.info("Measuring prompt processing...")
404            inputs, outputs = prepare_model_for_inference(args, model, config, tokenizer, prompt_length, prompt)
405            accelerator_prompt_latency_s, outputs = run_inference(args, model, args.num_runs, inputs, outputs)
406
407            # Calculate prompt metrics
408            accelerator_prompt_latency_ms = accelerator_prompt_latency_s * 1000
409            accelerator_prompt_thrpt = batch_size * (prompt_length / accelerator_prompt_latency_s)
410            logger.info(f"Average Latency of Prompt Processing: {accelerator_prompt_latency_ms} ms")
411            logger.info(
412                f"Average Throughput of Prompt Processing: {batch_size * (prompt_length / accelerator_prompt_latency_s)} tps"
413            )
414            csv_metrics.extend([accelerator_prompt_latency_ms, accelerator_prompt_thrpt])
415
416            # Measure token generation
417            logger.info("Measuring token generation...")
418            clear_cache()
419            inputs, outputs = prepare_model_for_inference(args, model, config, tokenizer, prompt_length, prompt)
420
421            all_token_ids = inputs["input_ids"].clone()
422            current_length = all_token_ids.shape[-1]
423            num_heads = config.num_key_value_heads
424            head_size = (
425                config.head_dim if hasattr(config, "head_dim") else config.hidden_size // config.num_attention_heads
426            )
427
428            has_eos = torch.zeros(batch_size, device=args.target_device, dtype=torch.bool)
429
430            # 0th entry will have prompt accelerator time, 1st entry onwards will have token generation accelerator time
431            accelerator_times = []
432            sampling_times = []  # cost to sample after each model run
433
434            wall_clock_start_time = time.perf_counter()
435            while current_length <= max_length:
436                # Run inference
437                accelerator_time_latency_s, outputs = run_inference(args, model, 1, inputs, outputs)
438                accelerator_times.append(accelerator_time_latency_s)
439
440                # Sample with argmax (greedy search)
441                sampling_start_time = time.perf_counter()
442                if outputs["logits"].shape[1] > 1:
443                    prompt_end_indices = inputs["attention_mask"].sum(1) - 1
444                    idxs = (
445                        prompt_end_indices.unsqueeze(dim=1)
446                        .repeat(1, config.vocab_size)
447                        .view(batch_size, 1, config.vocab_size)
448                    )
449                    next_token_logits = torch.gather(outputs["logits"], 1, idxs).squeeze()
450                else:
451                    next_token_logits = outputs["logits"][:, -1, :]
452                next_tokens = torch.argmax(next_token_logits, dim=-1)
453
454                # Check if we previously reached EOS token id or if generated token id is EOS token id
455                has_eos = has_eos | next_tokens == tokenizer.eos_token_id
456
457                # Determine which new tokens to add to list of all token ids
458                # Add EOS token ids for batch entries that ended early (ragged batching scenario where some batch entries ended early and some haven't)
459                tokens_to_add = next_tokens.masked_fill(has_eos, tokenizer.eos_token_id).reshape([batch_size, 1])
460                sampling_end_time = time.perf_counter()
461                sampling_times.append(sampling_end_time - sampling_start_time)
462
463                all_token_ids = torch.cat([all_token_ids, tokens_to_add], dim=-1)
464                current_length += 1
465
466                # Update inputs for next inference run
467                inputs["input_ids"] = tokens_to_add
468                inputs["attention_mask"] = torch.cat(
469                    [inputs["attention_mask"], (~has_eos).to(torch.int64).reshape(batch_size, 1)], 1
470                )
471                if "position_ids" in inputs:
472                    inputs["position_ids"] = torch.max(inputs["position_ids"], dim=1)[0].reshape(batch_size, 1) + 1
473
474                # Set logits to zeros for next inference run and re-use memory buffer
475                if outputs["logits"].shape[1] != 1:
476                    outputs["logits"] = outputs["logits"][:, :1, :].contiguous()
477                outputs["logits"].zero_()
478
479                # Update KV caches for next inference run
480                if args.engine == "pt":
481                    # Update KV caches for PyTorch
482                    inputs["past_key_values"] = outputs["past_key_values"]
483                elif not args.use_buffer_share:
484                    # Update KV caches for ONNX Runtime if buffer sharing is not used
485                    for i in range(config.num_hidden_layers):
486                        inputs[f"past_key_values.{i}.key"] = outputs[f"present.{i}.key"]
487                        inputs[f"past_key_values.{i}.value"] = outputs[f"present.{i}.value"]
488
489                    new_sequence_length = inputs["attention_mask"].shape[1]
490                    for i in range(config.num_hidden_layers):
491                        present_key = torch.zeros(
492                            batch_size,
493                            num_heads,
494                            new_sequence_length,
495                            head_size,
496                            device=args.target_device,
497                            dtype=args.torch_dtype,
498                        )
499                        present_value = torch.zeros(
500                            batch_size,
501                            num_heads,
502                            new_sequence_length,
503                            head_size,
504                            device=args.target_device,
505                            dtype=args.torch_dtype,
506                        )
507                        outputs.update(
508                            {
509                                f"present.{i}.key": present_key.contiguous(),
510                                f"present.{i}.value": present_value.contiguous(),
511                            }
512                        )
513
514            wall_clock_end_time = time.perf_counter()
515
516            # Filter out any anomaly accelerator times (e.g. for `torch.compile`)
517            accelerator_times.pop(0)  # Remove prompt processing time
518            if args.anomaly_filtering:
519                anomaly_threshold_factor = 10
520                min_time_s = min(accelerator_times)
521                orig_size = len(accelerator_times)
522                accelerator_times = list(
523                    filter(lambda acc_time: acc_time < anomaly_threshold_factor * min_time_s, accelerator_times)
524                )
525                new_size = len(accelerator_times)
526                logger.info(
527                    f"Filtered out {orig_size - new_size} anomaly accelerator times that are {anomaly_threshold_factor}x greater than {min_time_s * 1000} ms..."
528                )
529
530            #######################################################
531            # Calculate sampling and first token generated metrics
532            #######################################################
533
534            # Calculate sampling metrics
535            avg_sampling_latency_s = sum(sampling_times) / len(sampling_times)
536            avg_sampling_latency_ms = avg_sampling_latency_s * 1000
537            avg_sampling_thrpt = batch_size * (1 / avg_sampling_latency_s)
538            logger.info(f"Average Latency of Sampling: {avg_sampling_latency_ms} ms")
539            logger.info(f"Average Throughput of Sampling: {avg_sampling_thrpt} tps")
540
541            # Calculate first token generated metrics
542            first_token_latency_s = accelerator_times[0]
543            first_token_latency_ms = first_token_latency_s * 1000
544            first_token_thrpt = batch_size * (1 / first_token_latency_s)
545            logger.info(f"Latency of First Token Generated: {first_token_latency_ms} ms")
546            logger.info(f"Throughput of First Token Generated: {first_token_thrpt} tps")
547
548            ####################################################
549            # Calculate first `halfway` token generated metrics
550            ####################################################
551
552            halfway = args.generation_length // 2
553            halfway_token_latency_s = sum(accelerator_times[:halfway]) / len(accelerator_times[:halfway])
554            halfway_token_latency_ms = halfway_token_latency_s * 1000
555            halfway_token_thrpt = batch_size * (1 / halfway_token_latency_s)
556            logger.info(f"Average Latency of First {halfway} Tokens Generated: {halfway_token_latency_ms} ms")
557            logger.info(f"Average Throughput of First {halfway} Tokens Generated: {halfway_token_thrpt} tps")
558
559            #########################################
560            # Calculate all tokens generated metrics
561            #########################################
562
563            all_token_latency_s = sum(accelerator_times) / len(accelerator_times)
564            all_token_latency_ms = all_token_latency_s * 1000
565            all_token_thrpt = batch_size * (1 / all_token_latency_s)
566            logger.info(
567                f"Average Latency of First {args.generation_length} Tokens Generated: {all_token_latency_ms} ms"
568            )
569            logger.info(f"Average Throughput of First {args.generation_length} Tokens Generated: {all_token_thrpt} tps")
570
571            ###############################
572            # Calculate wall clock metrics
573            ###############################
574
575            wall_clock_latency_s = wall_clock_end_time - wall_clock_start_time
576            wall_clock_thrpt = batch_size * ((prompt_length + args.generation_length) / wall_clock_latency_s)
577            logger.info(f"Wall-Clock Latency: {wall_clock_latency_s} s")
578            logger.info(
579                f"Wall-Clock Throughput: {batch_size * ((prompt_length + args.generation_length) / wall_clock_latency_s)} tps"
580            )
581
582            # Add metrics to CSV
583            logger.info("Adding results to CSV")
584            csv_metrics.extend(
585                [
586                    avg_sampling_latency_ms,
587                    avg_sampling_thrpt,
588                    first_token_latency_ms,
589                    first_token_thrpt,
590                    halfway_token_latency_ms,
591                    halfway_token_thrpt,
592                    all_token_latency_ms,
593                    all_token_thrpt,
594                    wall_clock_latency_s,
595                    wall_clock_thrpt,
596                ]
597            )
598            all_csv_metrics.append(csv_metrics)
599
600        except Exception as e:
601            logger.info(f"Could not benchmark at batch size = {batch_size}, prompt length = {prompt_length} - {e}")
602
603    filename = f"benchmark_{args.engine}_e2e_{datetime.datetime.now():%Y-%m-%d_%H:%M:%S}.csv"
604    save_results(all_csv_metrics, filename, args.generation_length)
605
606
607if __name__ == "__main__":
608    main()
609 
codekingpro/portable-devtools · Team Ai