Team Ai
Datasetpublic

codekingpro/portable-devtools

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