codekingpro/portable-devtools
114k
1# -------------------------------------------------------------------------
2# Copyright (R) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5import logging
6import os
7import sys
8from collections.abc import Mapping
9
10import torch
11from sam2.build_sam import build_sam2
12from sam2.modeling.sam2_base import SAM2Base
13
14logger = logging.getLogger(__name__)
15
16
17def _get_model_cfg(model_type) -> str:
18 assert model_type in ["sam2_hiera_tiny", "sam2_hiera_small", "sam2_hiera_large", "sam2_hiera_base_plus"]
19 if model_type == "sam2_hiera_tiny":
20 model_cfg = "sam2_hiera_t.yaml"
21 elif model_type == "sam2_hiera_small":
22 model_cfg = "sam2_hiera_s.yaml"
23 elif model_type == "sam2_hiera_base_plus":
24 model_cfg = "sam2_hiera_b+.yaml"
25 else:
26 model_cfg = "sam2_hiera_l.yaml"
27 return model_cfg
28
29
30def load_sam2_model(sam2_dir, model_type, device: str | torch.device = "cpu") -> SAM2Base:
31 checkpoints_dir = os.path.join(sam2_dir, "checkpoints")
32 sam2_config_dir = os.path.join(sam2_dir, "sam2_configs")
33 if not os.path.exists(sam2_dir):
34 raise FileNotFoundError(f"{sam2_dir} does not exist. Please specify --sam2_dir correctly.")
35
36 if not os.path.exists(checkpoints_dir):
37 raise FileNotFoundError(f"{checkpoints_dir} does not exist. Please specify --sam2_dir correctly.")
38
39 if not os.path.exists(sam2_config_dir):
40 raise FileNotFoundError(f"{sam2_config_dir} does not exist. Please specify --sam2_dir correctly.")
41
42 checkpoint_path = os.path.join(checkpoints_dir, f"{model_type}.pt")
43 if not os.path.exists(checkpoint_path):
44 raise FileNotFoundError(f"{checkpoint_path} does not exist. Please download checkpoints under the directory.")
45
46 if sam2_dir not in sys.path:
47 sys.path.append(sam2_dir)
48
49 model_cfg = _get_model_cfg(model_type)
50 sam2_model = build_sam2(model_cfg, checkpoint_path, device=device)
51 return sam2_model
52
53
54def sam2_onnx_path(output_dir, model_type, component, multimask_output=False, suffix=""):
55 if component == "image_encoder":
56 return os.path.join(output_dir, f"{model_type}_image_encoder{suffix}.onnx")
57 elif component == "mask_decoder":
58 return os.path.join(output_dir, f"{model_type}_mask_decoder{suffix}.onnx")
59 elif component == "prompt_encoder":
60 return os.path.join(output_dir, f"{model_type}_prompt_encoder{suffix}.onnx")
61 else:
62 assert component == "image_decoder"
63 return os.path.join(
64 output_dir, f"{model_type}_image_decoder" + ("_multi" if multimask_output else "") + f"{suffix}.onnx"
65 )
66
67
68def encoder_shape_dict(batch_size: int, height: int, width: int) -> Mapping[str, list[int]]:
69 assert height == 1024 and width == 1024, "Only 1024x1024 images are supported."
70 return {
71 "image": [batch_size, 3, height, width],
72 "image_features_0": [batch_size, 32, height // 4, width // 4],
73 "image_features_1": [batch_size, 64, height // 8, width // 8],
74 "image_embeddings": [batch_size, 256, height // 16, width // 16],
75 }
76
77
78def decoder_shape_dict(
79 original_image_height: int,
80 original_image_width: int,
81 num_labels: int = 1,
82 max_points: int = 16,
83 num_masks: int = 1,
84) -> dict:
85 height: int = 1024
86 width: int = 1024
87 return {
88 "image_features_0": [1, 32, height // 4, width // 4],
89 "image_features_1": [1, 64, height // 8, width // 8],
90 "image_embeddings": [1, 256, height // 16, width // 16],
91 "point_coords": [num_labels, max_points, 2],
92 "point_labels": [num_labels, max_points],
93 "input_masks": [num_labels, 1, height // 4, width // 4],
94 "has_input_masks": [num_labels],
95 "original_image_size": [2],
96 "masks": [num_labels, num_masks, original_image_height, original_image_width],
97 "iou_predictions": [num_labels, num_masks],
98 "low_res_masks": [num_labels, num_masks, height // 4, width // 4],
99 }
100
101
102def compare_tensors_with_tolerance(
103 name: str,
104 tensor1: torch.Tensor,
105 tensor2: torch.Tensor,
106 atol=5e-3,
107 rtol=1e-4,
108 mismatch_percentage_tolerance=0.1,
109) -> bool:
110 assert tensor1.shape == tensor2.shape
111 a = tensor1.clone().float()
112 b = tensor2.clone().float()
113
114 differences = torch.abs(a - b)
115 mismatch_count = (differences > (rtol * torch.max(torch.abs(a), torch.abs(b)) + atol)).sum().item()
116
117 total_elements = a.numel()
118 mismatch_percentage = (mismatch_count / total_elements) * 100
119
120 passed = mismatch_percentage < mismatch_percentage_tolerance
121
122 log_func = logger.error if not passed else logger.info
123 log_func(
124 "%s: mismatched elements percentage %.2f (%d/%d). Verification %s (threshold=%.2f).",
125 name,
126 mismatch_percentage,
127 mismatch_count,
128 total_elements,
129 "passed" if passed else "failed",
130 mismatch_percentage_tolerance,
131 )
132
133 return passed
134
135
136def random_sam2_input_image(batch_size=1, image_height=1024, image_width=1024) -> torch.Tensor:
137 image = torch.randn(batch_size, 3, image_height, image_width, dtype=torch.float32).cpu()
138 return image
139
140
141def setup_logger(verbose=True):
142 if verbose:
143 logging.basicConfig(format="[%(filename)s:%(lineno)s - %(funcName)20s()] %(message)s")
144 logging.getLogger().setLevel(logging.INFO)
145 else:
146 logging.basicConfig(format="[%(message)s")
147 logging.getLogger().setLevel(logging.WARNING)
148 