codekingpro/portable-devtools
114k
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 