Team Ai
Apppublic

Amordia/Interactive-Automatic-Image-Labeling-Platform-Development

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
predictor.py242 linesDownload Raw Back to root
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