Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
prompt_encoder.py190 linesDownload Raw Back to sam2
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 
codekingpro/portable-devtools · Team Ai