Amordia/Interactive-Automatic-Image-Labeling-Platform-Development
0
1import torch2import torch.nn.functional as F3from typing import Dict, Tuple, Optional4import network5 6class Predictor:7 """8 Wrapper for ScribblePrompt Unet model9 """10 def __init__(self, path: str, verbose: bool = False):11 12 self.verbose = verbose13 14 assert path.exists(), f"Checkpoint {path} does not exist"15 self.path = path16 17 self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")18 self.build_model()19 self.load()20 self.model.eval()21 self.to_device()22 23 def build_model(self):24 """25 Build the model26 """27 self.model = network.UNet(28 in_channels = 5,29 out_channels = 1,30 features = [192, 192, 192, 192],31 )32 33 def load(self):34 """35 Load the state of the model from a checkpoint file.36 """37 with (self.path).open("rb") as f:38 state = torch.load(f, map_location=self.device)39 self.model.load_state_dict(state, strict=True)40 if self.verbose:41 print(42 f"Loaded checkpoint from {self.path} to {self.device}"43 )44 45 def to_device(self):46 """47 Move the model to cpu or gpu48 """49 self.device = "cuda" if torch.cuda.is_available() else "cpu"50 self.model = self.model.to(self.device)51 52 def predict(self, prompts: Dict[str,any], img_features: Optional[torch.Tensor] = None, multimask_mode: bool = False):53 """54 Make predictions!55 56 Returns:57 mask (torch.Tensor): H x W58 img_features (torch.Tensor): B x 1 x H x W (for SAM models)59 low_res_mask (torch.Tensor): B x 1 x H x W logits60 """61 if self.verbose:62 print("point_coords", prompts.get("point_coords", None))63 print("point_labels", prompts.get("point_labels", None))64 print("box", prompts.get("box", None))65 print("img", prompts.get("img").shape, prompts.get("img").min(), prompts.get("img").max())66 if prompts.get("scribble") is not None:67 print("scribble", prompts.get("scribble", None).shape, prompts.get("scribble").min(), prompts.get("scribble").max())68 69 original_shape = prompts.get('img').shape[-2:]70 71 # Rescale to 128 x 12872 prompts = rescale_inputs(prompts)73 74 # Prepare inputs for ScribblePrompt unet (1 x 5 x 128 x 128)75 x = prepare_inputs(prompts).float()76 77 with torch.no_grad():78 yhat = self.model(x.to(self.device)).cpu()79 80 mask = torch.sigmoid(yhat)81 82 # Resize for app resolution83 mask = F.interpolate(mask, size=original_shape, mode='bilinear').squeeze()84 85 # mask: H x W, yhat: 1 x 1 x H x W86 return mask, None, yhat87 88 89# -----------------------------------------------------------------------------90# Prepare inputs91# -----------------------------------------------------------------------------92 93def rescale_inputs(inputs: Dict[str,any], res=128):94 """95 Rescale the inputs 96 """ 97 h,w = inputs['img'].shape[-2:]98 99 if h != res or w != res:100 101 inputs.update(dict(102 img = F.interpolate(inputs['img'], size=(res,res), mode='bilinear')103 ))104 105 if inputs.get('scribble') is not None:106 inputs.update({107 'scribble': F.interpolate(inputs['scribble'], size=(res,res), mode='bilinear') 108 })109 110 if inputs.get("box") is not None:111 boxes = inputs.get("box").clone()112 coords = boxes.reshape(-1, 2, 2)113 coords[..., 0] = coords[..., 0] * (res / w)114 coords[..., 1] = coords[..., 1] * (res / h)115 inputs.update({'box': coords.reshape(1, -1, 4).int()})116 117 if inputs.get("point_coords") is not None:118 coords = inputs.get("point_coords").clone()119 coords[..., 0] = coords[..., 0] * (res / w)120 coords[..., 1] = coords[..., 1] * (res / h)121 inputs.update({'point_coords': coords.int()})122 123 return inputs124 125def prepare_inputs(inputs: Dict[str,torch.Tensor], device = None) -> torch.Tensor:126 """127 Prepare inputs for ScribblePrompt Unet128 129 Returns: 130 x (torch.Tensor): B x 5 x H x W131 """132 img = inputs['img']133 if device is None:134 device = img.device135 136 img = img.to(device)137 shape = tuple(img.shape[-2:])138 139 if inputs.get("box") is not None:140 # Embed bounding box141 # Input: B x 1 x 4 142 # Output: B x 1 x H x W143 box_embed = bbox_shaded(inputs['box'], shape=shape, device=device)144 else:145 box_embed = torch.zeros(img.shape, device=device)146 147 if inputs.get("point_coords") is not None:148 # Encode points149 # B x 2 x H x W150 scribble_click_embed = click_onehot(inputs['point_coords'], inputs['point_labels'], shape=shape)151 else:152 scribble_click_embed = torch.zeros((img.shape[0], 2) + shape, device=device)153 154 if inputs.get("scribble") is not None:155 # Combine scribbles with click encoding156 # B x 2 x H x W157 scribble_click_embed = torch.clamp(scribble_click_embed + inputs.get('scribble'), min=0.0, max=1.0)158 159 if inputs.get('mask_input') is not None:160 # Previous prediction161 mask_input = inputs['mask_input']162 else:163 # Initialize empty channel for mask input164 mask_input = torch.zeros(img.shape, device=img.device)165 166 x = torch.cat((img, box_embed, scribble_click_embed, mask_input), dim=-3)167 # B x 5 x H x W168 169 return x170 171# -----------------------------------------------------------------------------172# Encode clicks and bounding boxes173# -----------------------------------------------------------------------------174 175def click_onehot(point_coords, point_labels, shape: Tuple[int,int] = (128,128), indexing='xy'):176 """177 Represent clicks as two HxW binary masks (one for positive clicks and one for negative) 178 with 1 at the click locations and 0 otherwise179 180 Args:181 point_coords (torch.Tensor): BxNx2 tensor of xy coordinates182 point_labels (torch.Tensor): BxN tensor of labels (0 or 1)183 shape (tuple): output shape 184 Returns:185 embed (torch.Tensor): Bx2xHxW tensor 186 """187 assert indexing in ['xy','uv'], f"Invalid indexing: {indexing}"188 assert len(point_coords.shape) == 3, "point_coords must be BxNx2"189 assert point_coords.shape[-1] == 2, "point_coords must be BxNx2"190 assert point_labels.shape[-1] == point_coords.shape[1], "point_labels must be BxN"191 assert len(shape)==2, f"shape must be 2D: {shape}"192 193 device = point_coords.device194 batch_size = point_coords.shape[0]195 n_points = point_coords.shape[1]196 197 embed = torch.zeros((batch_size,2)+shape, device=device)198 labels = point_labels.flatten().float()199 200 idx_coords = torch.cat((201 torch.arange(batch_size, device=device).reshape(-1,1).repeat(1,n_points)[...,None], 202 point_coords203 ), axis=2).reshape(-1,3)204 205 if indexing=='xy':206 embed[ idx_coords[:,0], 0, idx_coords[:,2], idx_coords[:,1] ] = labels207 embed[ idx_coords[:,0], 1, idx_coords[:,2], idx_coords[:,1] ] = 1.0-labels208 else:209 embed[ idx_coords[:,0], 0, idx_coords[:,1], idx_coords[:,2] ] = labels210 embed[ idx_coords[:,0], 1, idx_coords[:,1], idx_coords[:,2] ] = 1.0-labels211 212 return embed213 214 215def bbox_shaded(boxes, shape: Tuple[int,int] = (128,128), device='cpu'):216 """217 Represent bounding boxes as a binary mask with 1 inside boxes and 0 otherwise218 219 Args:220 boxes (torch.Tensor): Bx1x4 [x1, y1, x2, y2]221 Returns:222 bbox_embed (torch.Tesor): Bx1xHxW according to shape223 """224 assert len(shape)==2, "shape must be 2D"225 if isinstance(boxes, torch.Tensor):226 boxes = boxes.int().cpu().numpy()227 228 batch_size = boxes.shape[0]229 n_boxes = boxes.shape[1]230 bbox_embed = torch.zeros((batch_size,1)+tuple(shape), device=device, dtype=torch.float32)231 232 if boxes is not None:233 for i in range(batch_size):234 for j in range(n_boxes):235 x1, y1, x2, y2 = boxes[i,j,:]236 x_min = min(x1,x2)237 x_max = max(x1,x2)238 y_min = min(y1,y2)239 y_max = max(y1,y2)240 bbox_embed[ i, 0, y_min:y_max, x_min:x_max ] = 1.0241 242 return bbox_embed