codekingpro/portable-devtools
114k
1# -------------------------------------------------------------------------
2# Copyright (R) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5import logging
6
7import torch
8from sam2.modeling.sam2_base import SAM2Base
9from sam2_utils import compare_tensors_with_tolerance
10from torch import nn
11
12logger = logging.getLogger(__name__)
13
14
15class SAM2PromptEncoder(nn.Module):
16 def __init__(self, sam_model: SAM2Base):
17 super().__init__()
18 self.prompt_encoder = sam_model.sam_prompt_encoder
19 self.model = sam_model
20
21 @torch.no_grad()
22 def forward(
23 self,
24 point_coords: torch.Tensor,
25 point_labels: torch.Tensor,
26 input_masks: torch.Tensor,
27 has_input_masks: torch.Tensor,
28 ):
29 """Encode prompts.
30
31 Args:
32 point_coords (torch.Tensor): [L, P, 2] shape and float32 dtype and contains the absolute pixel
33 coordinate in (x, y) format of the P input points in image of size 1024x1024.
34 point_labels (torch.Tensor): shape [L, P] and int32 dtype, where 1 means
35 positive (foreground), 0 means negative (background), -1 means padding,
36 2 (box left upper corner), 3 (box right bottom corner).
37 input_masks (torch.Tensor): [L, 1, H/4, W/4]. Low resolution mask input to the model.
38 Typically coming from a previous iteration.
39 has_input_masks (torch.Tensor): [L]. 1.0 if input_masks is used, 0.0 otherwise.
40 Returns:
41 sparse_embeddings (torch.Tensor): [L, P+1, 256], embedding for points and boxes.
42 dense_embeddings (torch.Tensor): [L, 256, 64, 64]. embedding for input masks.
43 image_pe (torch.Tensor, optional): [1, 256, 64, 64]. image positional encoding.
44 """
45 sparse_embeddings = self._embed_points(point_coords, point_labels)
46 dense_embeddings = self._embed_masks(input_masks, has_input_masks)
47 image_pe = self.prompt_encoder.get_dense_pe()
48
49 return sparse_embeddings, dense_embeddings, image_pe
50
51 def _embed_points(self, point_coords: torch.Tensor, point_labels: torch.Tensor) -> torch.Tensor:
52 point_coords = point_coords + 0.5
53
54 padding_point = torch.zeros((point_coords.shape[0], 1, 2), device=point_coords.device)
55 padding_label = -torch.ones((point_labels.shape[0], 1), device=point_labels.device)
56 point_coords = torch.cat([point_coords, padding_point], dim=1)
57 point_labels = torch.cat([point_labels, padding_label], dim=1)
58
59 # Note that the input coordinates are based on image size 1024x1024. Here we normalize it to [0.0, 1.0).
60 point_coords[:, :, 0] = point_coords[:, :, 0] / self.model.image_size
61 point_coords[:, :, 1] = point_coords[:, :, 1] / self.model.image_size
62
63 point_embedding = self.prompt_encoder.pe_layer._pe_encoding(point_coords)
64 point_labels = point_labels.unsqueeze(-1).expand_as(point_embedding)
65
66 point_embedding = point_embedding * (point_labels != -1)
67 point_embedding = point_embedding + self.prompt_encoder.not_a_point_embed.weight * (point_labels == -1)
68
69 for i in range(self.prompt_encoder.num_point_embeddings):
70 point_embedding = point_embedding + self.prompt_encoder.point_embeddings[i].weight * (point_labels == i)
71
72 return point_embedding
73
74 def _embed_masks(self, input_masks: torch.Tensor, has_input_masks: torch.Tensor) -> torch.Tensor:
75 mask_embedding = self.prompt_encoder.mask_downscaling(input_masks)
76 no_mask_embedding = self.prompt_encoder.no_mask_embed.weight.reshape(1, -1, 1, 1)
77 logger.info("no_mask_embedding.shape: %s", no_mask_embedding.shape)
78 mask_embedding = has_input_masks * mask_embedding + (1.0 - has_input_masks) * no_mask_embedding
79 logger.info("mask_embedding.shape: %s", mask_embedding.shape)
80 return mask_embedding
81
82
83def export_prompt_encoder_onnx(
84 sam2_model: SAM2Base,
85 onnx_model_path: str,
86):
87 sam2_prompt_encoder = SAM2PromptEncoder(sam2_model).cpu()
88
89 num_labels = 2
90 num_points = 3
91 point_coords = torch.randint(low=0, high=1024, size=(num_labels, num_points, 2), dtype=torch.float)
92 point_labels = torch.randint(low=0, high=1, size=(num_labels, num_points), dtype=torch.int32)
93 input_masks = torch.zeros(num_labels, 1, 256, 256, dtype=torch.float)
94 has_input_masks = torch.ones(1, dtype=torch.float)
95
96 sparse_embeddings, dense_embeddings, image_pe = sam2_prompt_encoder(
97 point_coords, point_labels, input_masks, has_input_masks
98 )
99
100 logger.info("point_coords.shape: %s", point_coords.shape)
101 logger.info("point_labels.shape: %s", point_labels.shape)
102 logger.info("input_masks.shape: %s", input_masks.shape)
103 logger.info("has_input_masks.shape: %s", has_input_masks.shape)
104
105 logger.info("sparse_embeddings.shape: %s", sparse_embeddings.shape)
106 logger.info("dense_embeddings.shape: %s", dense_embeddings.shape)
107 logger.info("image_pe.shape: %s", image_pe.shape)
108
109 torch.onnx.export(
110 sam2_prompt_encoder,
111 (point_coords, point_labels, input_masks, has_input_masks),
112 onnx_model_path,
113 export_params=True,
114 opset_version=18,
115 do_constant_folding=True,
116 input_names=["point_coords", "point_labels", "input_masks", "has_input_masks"],
117 output_names=["sparse_embeddings", "dense_embeddings", "image_pe"],
118 dynamic_axes={
119 "point_coords": {0: "num_labels", 1: "num_points"},
120 "point_labels": {0: "num_labels", 1: "num_points"},
121 "input_masks": {0: "num_labels"},
122 "sparse_embeddings": {0: "num_labels", 1: "num_points+1"},
123 "dense_embeddings": {0: "num_labels"},
124 },
125 )
126
127 print("prompt encoder onnx model saved to ", onnx_model_path)
128
129
130def test_prompt_encoder_onnx(
131 sam2_model: SAM2Base,
132 onnx_model_path: str,
133):
134 sam2_prompt_encoder = SAM2PromptEncoder(sam2_model).cpu()
135
136 num_labels = 1
137 num_points = 5
138 point_coords = torch.randint(low=0, high=1024, size=(num_labels, num_points, 2), dtype=torch.float)
139 point_labels = torch.randint(low=0, high=1, size=(num_labels, num_points), dtype=torch.int32)
140 input_masks = torch.rand(num_labels, 1, 256, 256, dtype=torch.float)
141 has_input_masks = torch.ones(1, dtype=torch.float)
142
143 sparse_embeddings, dense_embeddings, image_pe = sam2_prompt_encoder(
144 point_coords, point_labels, input_masks, has_input_masks
145 )
146
147 import onnxruntime # noqa: PLC0415
148
149 ort_session = onnxruntime.InferenceSession(onnx_model_path, providers=["CPUExecutionProvider"])
150
151 model_inputs = ort_session.get_inputs()
152 input_names = [model_inputs[i].name for i in range(len(model_inputs))]
153 logger.info("input_names: %s", input_names)
154
155 model_outputs = ort_session.get_outputs()
156 output_names = [model_outputs[i].name for i in range(len(model_outputs))]
157 logger.info("output_names: %s", output_names)
158
159 outputs = ort_session.run(
160 output_names,
161 {
162 "point_coords": point_coords.numpy(),
163 "point_labels": point_labels.numpy(),
164 "input_masks": input_masks.numpy(),
165 "has_input_masks": has_input_masks.numpy(),
166 },
167 )
168
169 for i, output_name in enumerate(output_names):
170 logger.info("output %s shape: %s", output_name, outputs[i].shape)
171
172 ort_sparse_embeddings, ort_dense_embeddings, ort_image_pe = outputs
173 if (
174 compare_tensors_with_tolerance(
175 "sparse_embeddings",
176 sparse_embeddings,
177 torch.tensor(ort_sparse_embeddings),
178 mismatch_percentage_tolerance=0.2,
179 )
180 and compare_tensors_with_tolerance(
181 "dense_embeddings", dense_embeddings, torch.tensor(ort_dense_embeddings), mismatch_percentage_tolerance=0.2
182 )
183 and compare_tensors_with_tolerance(
184 "image_pe", image_pe, torch.tensor(ort_image_pe), mismatch_percentage_tolerance=0.2
185 )
186 ):
187 print(f"onnx model has been verified: {onnx_model_path}")
188 else:
189 print(f"onnx model verification failed: {onnx_model_path}")
190 