Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
benchmark_sam2.py639 linesDownload Raw Back to sam2
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation.  All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5
6"""
7Benchmark performance of SAM2 encoder with ORT or PyTorch. See benchmark_sam2.sh for usage.
8"""
9
10import argparse
11import csv
12import statistics
13import time
14from collections.abc import Mapping
15from datetime import datetime
16
17import torch
18from image_decoder import SAM2ImageDecoder
19from image_encoder import SAM2ImageEncoder
20from sam2_utils import decoder_shape_dict, encoder_shape_dict, load_sam2_model
21
22from onnxruntime import InferenceSession, SessionOptions, get_available_providers
23from onnxruntime.transformers.io_binding_helper import CudaSession
24
25
26class TestConfig:
27    def __init__(
28        self,
29        model_type: str,
30        onnx_path: str,
31        sam2_dir: str,
32        device: torch.device,
33        component: str = "image_encoder",
34        provider="CPUExecutionProvider",
35        torch_compile_mode="max-autotune",
36        batch_size: int = 1,
37        height: int = 1024,
38        width: int = 1024,
39        num_labels: int = 1,
40        num_points: int = 1,
41        num_masks: int = 1,
42        multi_mask_output: bool = False,
43        use_tf32: bool = True,
44        enable_cuda_graph: bool = False,
45        dtype=torch.float32,
46        prefer_nhwc: bool = False,
47        warm_up: int = 5,
48        enable_nvtx_profile: bool = False,
49        enable_ort_profile: bool = False,
50        enable_torch_profile: bool = False,
51        repeats: int = 1000,
52        verbose: bool = False,
53    ):
54        assert model_type in ["sam2_hiera_tiny", "sam2_hiera_small", "sam2_hiera_large", "sam2_hiera_base_plus"]
55        assert height >= 160 and height <= 4096
56        assert width >= 160 and width <= 4096
57
58        self.model_type = model_type
59        self.onnx_path = onnx_path
60        self.sam2_dir = sam2_dir
61        self.component = component
62        self.provider = provider
63        self.torch_compile_mode = torch_compile_mode
64        self.batch_size = batch_size
65        self.height = height
66        self.width = width
67        self.num_labels = num_labels
68        self.num_points = num_points
69        self.num_masks = num_masks
70        self.multi_mask_output = multi_mask_output
71        self.device = device
72        self.use_tf32 = use_tf32
73        self.enable_cuda_graph = enable_cuda_graph
74        self.dtype = dtype
75        self.prefer_nhwc = prefer_nhwc
76        self.warm_up = warm_up
77        self.enable_nvtx_profile = enable_nvtx_profile
78        self.enable_ort_profile = enable_ort_profile
79        self.enable_torch_profile = enable_torch_profile
80        self.repeats = repeats
81        self.verbose = verbose
82
83        if self.component == "image_encoder":
84            assert self.height == 1024 and self.width == 1024, "Only image size 1024x1024 is allowed for image encoder."
85
86    def __repr__(self):
87        return f"{vars(self)}"
88
89    def shape_dict(self) -> Mapping[str, list[int]]:
90        if self.component == "image_encoder":
91            return encoder_shape_dict(self.batch_size, self.height, self.width)
92        else:
93            return decoder_shape_dict(self.height, self.width, self.num_labels, self.num_points, self.num_masks)
94
95    def random_inputs(self) -> Mapping[str, torch.Tensor]:
96        dtype = self.dtype
97        if self.component == "image_encoder":
98            return {"image": torch.randn(self.batch_size, 3, self.height, self.width, dtype=dtype, device=self.device)}
99        else:
100            return {
101                "image_features_0": torch.rand(1, 32, 256, 256, dtype=dtype, device=self.device),
102                "image_features_1": torch.rand(1, 64, 128, 128, dtype=dtype, device=self.device),
103                "image_embeddings": torch.rand(1, 256, 64, 64, dtype=dtype, device=self.device),
104                "point_coords": torch.randint(
105                    0, 1024, (self.num_labels, self.num_points, 2), dtype=dtype, device=self.device
106                ),
107                "point_labels": torch.randint(
108                    0, 1, (self.num_labels, self.num_points), dtype=torch.int32, device=self.device
109                ),
110                "input_masks": torch.zeros(self.num_labels, 1, 256, 256, dtype=dtype, device=self.device),
111                "has_input_masks": torch.ones(self.num_labels, dtype=dtype, device=self.device),
112                "original_image_size": torch.tensor([self.height, self.width], dtype=torch.int32, device=self.device),
113            }
114
115
116def create_ort_session(config: TestConfig, session_options=None) -> InferenceSession:
117    if config.verbose:
118        print(f"create session for {vars(config)}")
119
120    if config.provider == "CUDAExecutionProvider":
121        device_id = torch.cuda.current_device() if isinstance(config.device, str) else config.device.index
122        provider_options = CudaSession.get_cuda_provider_options(device_id, config.enable_cuda_graph)
123        provider_options["use_tf32"] = int(config.use_tf32)
124        if config.prefer_nhwc:
125            provider_options["prefer_nhwc"] = 1
126        providers = [(config.provider, provider_options), "CPUExecutionProvider"]
127    else:
128        providers = ["CPUExecutionProvider"]
129
130    ort_session = InferenceSession(config.onnx_path, session_options, providers=providers)
131    return ort_session
132
133
134def create_session(config: TestConfig, session_options=None) -> CudaSession:
135    ort_session = create_ort_session(config, session_options)
136    cuda_session = CudaSession(ort_session, config.device, config.enable_cuda_graph)
137    cuda_session.allocate_buffers(config.shape_dict())
138    return cuda_session
139
140
141class OrtTestSession:
142    """A wrapper of ORT session to test relevance and performance."""
143
144    def __init__(self, config: TestConfig, session_options=None):
145        self.ort_session = create_session(config, session_options)
146        self.feed_dict = config.random_inputs()
147
148    def infer(self):
149        return self.ort_session.infer(self.feed_dict)
150
151
152def measure_latency(cuda_session: CudaSession, input_dict):
153    start = time.time()
154    _ = cuda_session.infer(input_dict)
155    end = time.time()
156    return end - start
157
158
159def run_torch(config: TestConfig):
160    device_type = config.device.type
161    is_cuda = device_type == "cuda"
162
163    # Turn on TF32 for Ampere GPUs which could help when data type is float32.
164    if is_cuda and torch.cuda.get_device_properties(0).major >= 8 and config.use_tf32:
165        torch.backends.cuda.matmul.allow_tf32 = True
166        torch.backends.cudnn.allow_tf32 = True
167
168    enabled_auto_cast = is_cuda and config.dtype != torch.float32
169    ort_inputs = config.random_inputs()
170
171    with torch.inference_mode(), torch.autocast(device_type=device_type, dtype=config.dtype, enabled=enabled_auto_cast):
172        sam2_model = load_sam2_model(config.sam2_dir, config.model_type, device=config.device)
173        if config.component == "image_encoder":
174            if is_cuda and config.torch_compile_mode != "none":
175                sam2_model.image_encoder.forward = torch.compile(
176                    sam2_model.image_encoder.forward,
177                    mode=config.torch_compile_mode,  # "reduce-overhead" if you want to reduce latency of first run.
178                    fullgraph=True,
179                    dynamic=False,
180                )
181
182            image_shape = config.shape_dict()["image"]
183            img = torch.randn(image_shape).to(device=config.device, dtype=config.dtype)
184            sam2_encoder = SAM2ImageEncoder(sam2_model)
185
186            if is_cuda and config.torch_compile_mode != "none":
187                print(f"Running warm up. It will take a while since torch compile mode is {config.torch_compile_mode}.")
188
189            for _ in range(config.warm_up):
190                _image_features_0, _image_features_1, _image_embeddings = sam2_encoder(img)
191
192            if is_cuda and config.enable_nvtx_profile:
193                import nvtx  # noqa: PLC0415
194                from cuda import cudart  # noqa: PLC0415
195
196                cudart.cudaProfilerStart()
197                print("Start nvtx profiling on encoder ...")
198                with nvtx.annotate("one_run"):
199                    sam2_encoder(img, enable_nvtx_profile=True)
200                cudart.cudaProfilerStop()
201
202            if is_cuda and config.enable_torch_profile:
203                with torch.profiler.profile(
204                    activities=[torch.profiler.ProfilerActivity.CPU, torch.profiler.ProfilerActivity.CUDA],
205                    record_shapes=True,
206                ) as prof:
207                    print("Start torch profiling on encoder ...")
208                    with torch.profiler.record_function("encoder"):
209                        sam2_encoder(img)
210                print(prof.key_averages().table(sort_by="cuda_time_total", row_limit=10))
211                prof.export_chrome_trace("torch_image_encoder.json")
212
213            if config.repeats == 0:
214                return
215
216            print(f"Start {config.repeats} runs of performance tests...")
217            start = time.time()
218            for _ in range(config.repeats):
219                _image_features_0, _image_features_1, _image_embeddings = sam2_encoder(img)
220                if is_cuda:
221                    torch.cuda.synchronize()
222        else:
223            torch_inputs = (
224                ort_inputs["image_features_0"],
225                ort_inputs["image_features_1"],
226                ort_inputs["image_embeddings"],
227                ort_inputs["point_coords"],
228                ort_inputs["point_labels"],
229                ort_inputs["input_masks"],
230                ort_inputs["has_input_masks"],
231                ort_inputs["original_image_size"],
232            )
233
234            sam2_decoder = SAM2ImageDecoder(
235                sam2_model,
236                multimask_output=config.multi_mask_output,
237            )
238
239            if is_cuda and config.torch_compile_mode != "none":
240                sam2_decoder.forward = torch.compile(
241                    sam2_decoder.forward,
242                    mode=config.torch_compile_mode,
243                    fullgraph=True,
244                    dynamic=False,
245                )
246
247            # warm up
248            for _ in range(config.warm_up):
249                _masks, _iou_predictions, _low_res_masks = sam2_decoder(*torch_inputs)
250
251            if is_cuda and config.enable_nvtx_profile:
252                import nvtx  # noqa: PLC0415
253                from cuda import cudart  # noqa: PLC0415
254
255                cudart.cudaProfilerStart()
256                print("Start nvtx profiling on decoder...")
257                with nvtx.annotate("one_run"):
258                    sam2_decoder(*torch_inputs, enable_nvtx_profile=True)
259                cudart.cudaProfilerStop()
260
261            if is_cuda and config.enable_torch_profile:
262                with torch.profiler.profile(
263                    activities=[torch.profiler.ProfilerActivity.CPU, torch.profiler.ProfilerActivity.CUDA],
264                    record_shapes=True,
265                ) as prof:
266                    print("Start torch profiling on decoder ...")
267                    with torch.profiler.record_function("decoder"):
268                        sam2_decoder(*torch_inputs)
269                print(prof.key_averages().table(sort_by="cuda_time_total", row_limit=10))
270                prof.export_chrome_trace("torch_image_decoder.json")
271
272            if config.repeats == 0:
273                return
274
275            print(f"Start {config.repeats} runs of performance tests...")
276            start = time.time()
277            for _ in range(config.repeats):
278                _masks, _iou_predictions, _low_res_masks = sam2_decoder(*torch_inputs)
279                if is_cuda:
280                    torch.cuda.synchronize()
281
282        end = time.time()
283        return (end - start) / config.repeats
284
285
286def run_test(
287    args: argparse.Namespace,
288    csv_writer: csv.DictWriter | None = None,
289):
290    use_gpu: bool = args.use_gpu
291    enable_cuda_graph: bool = args.use_cuda_graph
292    repeats: int = args.repeats
293
294    if use_gpu:
295        device_id = torch.cuda.current_device()
296        device = torch.device("cuda", device_id)
297        provider = "CUDAExecutionProvider"
298    else:
299        device_id = 0
300        device = torch.device("cpu")
301        enable_cuda_graph = False
302        provider = "CPUExecutionProvider"
303
304    dtypes = {"fp32": torch.float32, "fp16": torch.float16, "bf16": torch.bfloat16}
305    config = TestConfig(
306        model_type=args.model_type,
307        onnx_path=args.onnx_path,
308        sam2_dir=args.sam2_dir,
309        component=args.component,
310        provider=provider,
311        batch_size=args.batch_size,
312        height=args.height,
313        width=args.width,
314        device=device,
315        use_tf32=True,
316        enable_cuda_graph=enable_cuda_graph,
317        dtype=dtypes[args.dtype],
318        prefer_nhwc=args.prefer_nhwc,
319        repeats=args.repeats,
320        warm_up=args.warm_up,
321        enable_nvtx_profile=args.enable_nvtx_profile,
322        enable_ort_profile=args.enable_ort_profile,
323        enable_torch_profile=args.enable_torch_profile,
324        torch_compile_mode=args.torch_compile_mode,
325        verbose=False,
326    )
327
328    if args.engine == "ort":
329        sess_options = SessionOptions()
330        sess_options.intra_op_num_threads = args.intra_op_num_threads
331        if config.enable_ort_profile:
332            sess_options.enable_profiling = True
333            sess_options.log_severity_level = 4
334            sess_options.log_verbosity_level = 0
335
336        session = create_session(config, sess_options)
337        input_dict = config.random_inputs()
338
339        # warm up session
340        try:
341            for _ in range(config.warm_up):
342                _ = measure_latency(session, input_dict)
343        except Exception as e:
344            print(f"Failed to run {config=}. Exception: {e}")
345            return
346
347        if config.enable_nvtx_profile:
348            import nvtx  # noqa: PLC0415
349            from cuda import cudart  # noqa: PLC0415
350
351            cudart.cudaProfilerStart()
352            with nvtx.annotate("one_run"):
353                _ = session.infer(input_dict)
354            cudart.cudaProfilerStop()
355
356        if config.enable_ort_profile:
357            session.ort_session.end_profiling()
358
359        if repeats == 0:
360            return
361
362        latency_list = []
363        for _ in range(repeats):
364            latency = measure_latency(session, input_dict)
365            latency_list.append(latency)
366        average_latency = statistics.mean(latency_list)
367
368        del session
369    else:  # torch
370        with torch.no_grad():
371            try:
372                average_latency = run_torch(config)
373            except Exception as e:
374                print(f"Failed to run {config=}. Exception: {e}")
375                return
376
377        if repeats == 0:
378            return
379
380    engine = args.engine + ":" + ("cuda" if use_gpu else "cpu")
381    row = {
382        "model_type": args.model_type,
383        "component": args.component,
384        "dtype": args.dtype,
385        "use_gpu": use_gpu,
386        "enable_cuda_graph": enable_cuda_graph,
387        "prefer_nhwc": config.prefer_nhwc,
388        "use_tf32": config.use_tf32,
389        "batch_size": args.batch_size,
390        "height": args.height,
391        "width": args.width,
392        "multi_mask_output": args.multimask_output,
393        "num_labels": config.num_labels,
394        "num_points": config.num_points,
395        "num_masks": config.num_masks,
396        "intra_op_num_threads": args.intra_op_num_threads,
397        "warm_up": config.warm_up,
398        "repeats": repeats,
399        "enable_nvtx_profile": args.enable_nvtx_profile,
400        "torch_compile_mode": args.torch_compile_mode,
401        "engine": engine,
402        "average_latency": average_latency,
403    }
404
405    if csv_writer is not None:
406        csv_writer.writerow(row)
407
408    print(f"{vars(config)}")
409    print(f"{row}")
410
411
412def run_perf_test(args):
413    features = "gpu" if args.use_gpu else "cpu"
414    csv_filename = "benchmark_sam_{}_{}_{}.csv".format(
415        features,
416        args.engine,
417        datetime.now().strftime("%Y%m%d-%H%M%S"),
418    )
419    with open(csv_filename, mode="a", newline="") as csv_file:
420        column_names = [
421            "model_type",
422            "component",
423            "dtype",
424            "use_gpu",
425            "enable_cuda_graph",
426            "prefer_nhwc",
427            "use_tf32",
428            "batch_size",
429            "height",
430            "width",
431            "multi_mask_output",
432            "num_labels",
433            "num_points",
434            "num_masks",
435            "intra_op_num_threads",
436            "warm_up",
437            "repeats",
438            "enable_nvtx_profile",
439            "torch_compile_mode",
440            "engine",
441            "average_latency",
442        ]
443        csv_writer = csv.DictWriter(csv_file, fieldnames=column_names)
444        csv_writer.writeheader()
445
446        run_test(args, csv_writer)
447
448
449def _parse_arguments():
450    parser = argparse.ArgumentParser(description="Benchmark SMA2 for ONNX Runtime and PyTorch.")
451
452    parser.add_argument(
453        "--component",
454        required=False,
455        choices=["image_encoder", "image_decoder"],
456        default="image_encoder",
457        help="component to benchmark. Choices are image_encoder and image_decoder.",
458    )
459
460    parser.add_argument(
461        "--dtype", required=False, choices=["fp32", "fp16", "bf16"], default="fp32", help="Data type for inference."
462    )
463
464    parser.add_argument(
465        "--use_gpu",
466        required=False,
467        action="store_true",
468        help="Use GPU for inference.",
469    )
470    parser.set_defaults(use_gpu=False)
471
472    parser.add_argument(
473        "--use_cuda_graph",
474        required=False,
475        action="store_true",
476        help="Use cuda graph in onnxruntime.",
477    )
478    parser.set_defaults(use_cuda_graph=False)
479
480    parser.add_argument(
481        "--intra_op_num_threads",
482        required=False,
483        type=int,
484        choices=[0, 1, 2, 4, 8, 16],
485        default=0,
486        help="intra_op_num_threads for onnxruntime. ",
487    )
488
489    parser.add_argument(
490        "--batch_size",
491        required=False,
492        type=int,
493        default=1,
494        help="batch size",
495    )
496
497    parser.add_argument(
498        "--height",
499        required=False,
500        type=int,
501        default=1024,
502        help="image height",
503    )
504
505    parser.add_argument(
506        "--width",
507        required=False,
508        type=int,
509        default=1024,
510        help="image width",
511    )
512
513    parser.add_argument(
514        "--repeats",
515        required=False,
516        type=int,
517        default=1000,
518        help="number of repeats for performance test. Default is 1000.",
519    )
520
521    parser.add_argument(
522        "--warm_up",
523        required=False,
524        type=int,
525        default=5,
526        help="number of runs for warm up. Default is 5.",
527    )
528
529    parser.add_argument(
530        "--engine",
531        required=False,
532        type=str,
533        default="ort",
534        choices=["ort", "torch"],
535        help="engine for inference",
536    )
537
538    parser.add_argument(
539        "--multimask_output",
540        required=False,
541        default=False,
542        action="store_true",
543        help="Export mask_decoder or image_decoder with multimask_output",
544    )
545
546    parser.add_argument(
547        "--prefer_nhwc",
548        required=False,
549        default=False,
550        action="store_true",
551        help="Use prefer_nhwc=1 provider option for CUDAExecutionProvider",
552    )
553
554    parser.add_argument(
555        "--enable_nvtx_profile",
556        required=False,
557        default=False,
558        action="store_true",
559        help="Enable nvtx profiling. It will add an extra run for profiling before performance test.",
560    )
561
562    parser.add_argument(
563        "--enable_ort_profile",
564        required=False,
565        default=False,
566        action="store_true",
567        help="Enable ORT profiling.",
568    )
569
570    parser.add_argument(
571        "--enable_torch_profile",
572        required=False,
573        default=False,
574        action="store_true",
575        help="Enable PyTorch profiling. It will add an extra run for profiling before performance test.",
576    )
577
578    parser.add_argument(
579        "--model_type",
580        required=False,
581        type=str,
582        default="sam2_hiera_large",
583        choices=["sam2_hiera_tiny", "sam2_hiera_small", "sam2_hiera_large", "sam2_hiera_base_plus"],
584        help="sam2 model name",
585    )
586
587    parser.add_argument(
588        "--sam2_dir",
589        required=False,
590        type=str,
591        default="./segment-anything-2",
592        help="The directory of segment-anything-2 git root directory",
593    )
594
595    parser.add_argument(
596        "--onnx_path",
597        required=False,
598        type=str,
599        default="./sam2_onnx_models/sam2_hiera_large_image_encoder.onnx",
600        help="path of onnx model",
601    )
602
603    parser.add_argument(
604        "--torch_compile_mode",
605        required=False,
606        type=str,
607        default=None,
608        choices=["reduce-overhead", "max-autotune", "max-autotune-no-cudagraphs", "none"],
609        help="torch compile mode. none will disable torch compile.",
610    )
611
612    args = parser.parse_args()
613
614    return args
615
616
617if __name__ == "__main__":
618    args = _parse_arguments()
619    print(f"arguments:{args}")
620
621    if args.torch_compile_mode is None:
622        # image decoder will fail with compile modes other than "none".
623        args.torch_compile_mode = "max-autotune" if args.component == "image_encoder" else "none"
624
625    if args.use_gpu:
626        assert torch.cuda.is_available()
627        if args.engine == "ort":
628            assert "CUDAExecutionProvider" in get_available_providers()
629            args.enable_torch_profile = False
630    else:
631        # Only support cuda profiling for now.
632        assert not args.enable_nvtx_profile
633        assert not args.enable_torch_profile
634
635    if args.enable_nvtx_profile or args.enable_torch_profile:
636        run_test(args)
637    else:
638        run_perf_test(args)
639 
codekingpro/portable-devtools · Team Ai