# Copyright (c) Meta Platforms, Inc. and affiliates. # All rights reserved. # This source code is licensed under the license found in the # LICENSE file in the root directory of this source tree. import os import warnings from threading import Thread import numpy as np import torch from PIL import Image from tqdm import tqdm import subprocess import time from collections import defaultdict, deque import datetime import pickle import json import torch.distributed as dist import torch.nn.functional as F import cv2 def get_sdpa_settings(): if torch.cuda.is_available(): old_gpu = torch.cuda.get_device_properties(0).major < 7 # only use Flash Attention on Ampere (8.0) or newer GPUs use_flash_attn = torch.cuda.get_device_properties(0).major >= 8 if not use_flash_attn: warnings.warn( "Flash Attention is disabled as it requires a GPU with Ampere (8.0) CUDA capability.", category=UserWarning, stacklevel=2, ) # keep math kernel for PyTorch versions before 2.2 (Flash Attention v2 is only # available on PyTorch 2.2+, while Flash Attention v1 cannot handle all cases) pytorch_version = tuple(int(v) for v in torch.__version__.split(".")[:2]) if pytorch_version < (2, 2): warnings.warn( f"You are using PyTorch {torch.__version__} without Flash Attention v2 support. " "Consider upgrading to PyTorch 2.2+ for Flash Attention v2 (which could be faster).", category=UserWarning, stacklevel=2, ) math_kernel_on = pytorch_version < (2, 2) or not use_flash_attn else: old_gpu = True use_flash_attn = False math_kernel_on = True return old_gpu, use_flash_attn, math_kernel_on def get_connected_components(mask): """ Get the connected components (8-connectivity) of binary masks of shape (N, 1, H, W). Inputs: - mask: A binary mask tensor of shape (N, 1, H, W), where 1 is foreground and 0 is background. Outputs: - labels: A tensor of shape (N, 1, H, W) containing the connected component labels for foreground pixels and 0 for background pixels. - counts: A tensor of shape (N, 1, H, W) containing the area of the connected components for foreground pixels and 0 for background pixels. """ from src import _C return _C.get_connected_componnets(mask.to(torch.uint8).contiguous()) def mask_to_box(masks: torch.Tensor): """ compute bounding box given an input mask Inputs: - masks: [B, 1, H, W] masks, dtype=torch.Tensor Returns: - box_coords: [B, 1, 4], contains (x, y) coordinates of top left and bottom right box corners, dtype=torch.Tensor """ B, _, h, w = masks.shape device = masks.device xs = torch.arange(w, device=device, dtype=torch.int32) ys = torch.arange(h, device=device, dtype=torch.int32) grid_xs, grid_ys = torch.meshgrid(xs, ys, indexing="xy") grid_xs = grid_xs[None, None, ...].expand(B, 1, h, w) grid_ys = grid_ys[None, None, ...].expand(B, 1, h, w) min_xs, _ = torch.min(torch.where(masks, grid_xs, w).flatten(-2), dim=-1) max_xs, _ = torch.max(torch.where(masks, grid_xs, -1).flatten(-2), dim=-1) min_ys, _ = torch.min(torch.where(masks, grid_ys, h).flatten(-2), dim=-1) max_ys, _ = torch.max(torch.where(masks, grid_ys, -1).flatten(-2), dim=-1) bbox_coords = torch.stack((min_xs, min_ys, max_xs, max_ys), dim=-1) return bbox_coords def _load_img_as_tensor(img_path, image_size): img_pil = Image.open(img_path) img_np = np.array(img_pil.convert("RGB").resize((image_size, image_size))) if img_np.dtype == np.uint8: # np.uint8 is expected for JPEG images img_np = img_np / 255.0 else: raise RuntimeError(f"Unknown image dtype: {img_np.dtype} on {img_path}") img = torch.from_numpy(img_np).permute(2, 0, 1) video_width, video_height = img_pil.size # the original video size return img, video_height, video_width class AsyncVideoFrameLoader: """ A list of video frames to be load asynchronously without blocking session start. """ def __init__( self, img_paths, image_size, offload_video_to_cpu, img_mean, img_std, compute_device, ): self.img_paths = img_paths self.image_size = image_size self.offload_video_to_cpu = offload_video_to_cpu self.img_mean = img_mean self.img_std = img_std # items in `self.images` will be loaded asynchronously self.images = [None] * len(img_paths) # catch and raise any exceptions in the async loading thread self.exception = None # video_height and video_width be filled when loading the first image self.video_height = None self.video_width = None self.compute_device = compute_device # load the first frame to fill video_height and video_width and also # to cache it (since it's most likely where the user will click) self.__getitem__(0) # load the rest of frames asynchronously without blocking the session start def _load_frames(): try: for n in tqdm(range(len(self.images)), desc="frame loading (JPEG)"): self.__getitem__(n) except Exception as e: self.exception = e self.thread = Thread(target=_load_frames, daemon=True) self.thread.start() def __getitem__(self, index): if self.exception is not None: raise RuntimeError("Failure in frame loading thread") from self.exception img = self.images[index] if img is not None: return img img, video_height, video_width = _load_img_as_tensor( self.img_paths[index], self.image_size ) self.video_height = video_height self.video_width = video_width # normalize by mean and std img -= self.img_mean img /= self.img_std if not self.offload_video_to_cpu: img = img.to(self.compute_device, non_blocking=True) self.images[index] = img return img def __len__(self): return len(self.images) def load_video_frames( video_path, image_size, offload_video_to_cpu, img_mean=(0.485, 0.456, 0.406), img_std=(0.229, 0.224, 0.225), async_loading_frames=False, compute_device=torch.device("cuda"), ): """ Load the video frames from video_path. The frames are resized to image_size as in the model and are loaded to GPU if offload_video_to_cpu=False. This is used by the demo. """ is_bytes = isinstance(video_path, bytes) is_str = isinstance(video_path, str) is_mp4_path = is_str and os.path.splitext(video_path)[-1] in [".mp4", ".MP4"] if is_bytes or is_mp4_path: return load_video_frames_from_video_file( video_path=video_path, image_size=image_size, offload_video_to_cpu=offload_video_to_cpu, img_mean=img_mean, img_std=img_std, compute_device=compute_device, ) elif is_str and os.path.isdir(video_path): return load_video_frames_from_jpg_images( video_path=video_path, image_size=image_size, offload_video_to_cpu=offload_video_to_cpu, img_mean=img_mean, img_std=img_std, async_loading_frames=async_loading_frames, compute_device=compute_device, ) else: raise NotImplementedError( "Only MP4 video and JPEG folder are supported at this moment" ) def load_video_frames_from_jpg_images( video_path, image_size, offload_video_to_cpu, img_mean=(0.485, 0.456, 0.406), img_std=(0.229, 0.224, 0.225), async_loading_frames=False, compute_device=torch.device("cuda"), ): """ Load the video frames from a directory of JPEG files (".jpg" format). The frames are resized to image_size x image_size and are loaded to GPU if `offload_video_to_cpu` is `False` and to CPU if `offload_video_to_cpu` is `True`. You can load a frame asynchronously by setting `async_loading_frames` to `True`. """ if isinstance(video_path, str) and os.path.isdir(video_path): jpg_folder = video_path else: raise NotImplementedError( "Only JPEG frames are supported at this moment. For video files, you may use " "ffmpeg (https://ffmpeg.org/) to extract frames into a folder of JPEG files, such as \n" "```\n" "ffmpeg -i .mp4 -q:v 2 -start_number 0 /'%05d.jpg'\n" "```\n" "where `-q:v` generates high-quality JPEG frames and `-start_number 0` asks " "ffmpeg to start the JPEG file from 00000.jpg." ) frame_names = [ p for p in os.listdir(jpg_folder) if os.path.splitext(p)[-1] in [".jpg", ".jpeg", ".JPG", ".JPEG"] ] frame_names.sort(key=lambda p: int(os.path.splitext(p)[0])) num_frames = len(frame_names) if num_frames == 0: raise RuntimeError(f"no images found in {jpg_folder}") img_paths = [os.path.join(jpg_folder, frame_name) for frame_name in frame_names] img_mean = torch.tensor(img_mean, dtype=torch.float32)[:, None, None] img_std = torch.tensor(img_std, dtype=torch.float32)[:, None, None] if async_loading_frames: lazy_images = AsyncVideoFrameLoader( img_paths, image_size, offload_video_to_cpu, img_mean, img_std, compute_device, ) return lazy_images, lazy_images.video_height, lazy_images.video_width images = torch.zeros(num_frames, 3, image_size, image_size, dtype=torch.float32) for n, img_path in enumerate(tqdm(img_paths, desc="frame loading (JPEG)")): images[n], video_height, video_width = _load_img_as_tensor(img_path, image_size) if not offload_video_to_cpu: images = images.to(compute_device) img_mean = img_mean.to(compute_device) img_std = img_std.to(compute_device) # normalize by mean and std images -= img_mean images /= img_std return images, video_height, video_width def load_video_frames_from_video_file( video_path, image_size, offload_video_to_cpu, img_mean=(0.485, 0.456, 0.406), img_std=(0.229, 0.224, 0.225), compute_device=torch.device("cuda"), ): """Load the video frames from a video file.""" import decord img_mean = torch.tensor(img_mean, dtype=torch.float32)[:, None, None] img_std = torch.tensor(img_std, dtype=torch.float32)[:, None, None] # Get the original video height and width decord.bridge.set_bridge("torch") video_height, video_width, _ = decord.VideoReader(video_path).next().shape # Iterate over all frames in the video images = [] for frame in decord.VideoReader(video_path, width=image_size, height=image_size): images.append(frame.permute(2, 0, 1)) images = torch.stack(images, dim=0).float() / 255.0 if not offload_video_to_cpu: images = images.to(compute_device) img_mean = img_mean.to(compute_device) img_std = img_std.to(compute_device) # normalize by mean and std images -= img_mean images /= img_std return images, video_height, video_width def fill_holes_in_mask_scores(mask, max_area): """ A post processor to fill small holes in mask scores with area under `max_area`. """ # Holes are those connected components in background with area <= self.max_area # (background regions are those with mask scores <= 0) assert max_area > 0, "max_area must be positive" input_mask = mask try: labels, areas = get_connected_components(mask <= 0) is_hole = (labels > 0) & (areas <= max_area) # We fill holes with a small positive mask score (0.1) to change them to foreground. mask = torch.where(is_hole, 0.1, mask) except Exception as e: # Skip the post-processing step on removing small holes if the CUDA kernel fails warnings.warn( f"{e}\n\nSkipping the post-processing step due to the error above. You can " "still use SAM 2 and it's OK to ignore the error above, although some post-processing " "functionality may be limited (which doesn't affect the results in most cases; see " "https://github.com/facebookresearch/sam2/blob/main/INSTALL.md).", category=UserWarning, stacklevel=2, ) mask = input_mask return mask def concat_points(old_point_inputs, new_points, new_labels): """Add new points and labels to previous point inputs (add at the end).""" if old_point_inputs is None: points, labels = new_points, new_labels else: points = torch.cat([old_point_inputs["point_coords"], new_points], dim=1) labels = torch.cat([old_point_inputs["point_labels"], new_labels], dim=1) return {"point_coords": points, "point_labels": labels} class SmoothedValue(object): """Track a series of values and provide access to smoothed values over a window or the global series average. """ def __init__(self, window_size=20, fmt=None): if fmt is None: fmt = "{median:.4f} ({global_avg:.4f})" self.deque = deque(maxlen=window_size) self.total = 0.0 self.count = 0 self.fmt = fmt def update(self, value, n=1): self.deque.append(value) self.count += n self.total += value * n def synchronize_between_processes(self): """ Warning: does not synchronize the deque! """ if not is_dist_avail_and_initialized(): return t = torch.tensor([self.count, self.total], dtype=torch.float64, device='cuda') dist.barrier() dist.all_reduce(t) t = t.tolist() self.count = int(t[0]) self.total = t[1] @property def median(self): d = torch.tensor(list(self.deque)) if d.shape[0] == 0: return 0 return d.median().item() @property def avg(self): d = torch.tensor(list(self.deque), dtype=torch.float32) return d.mean().item() @property def global_avg(self): return self.total / self.count @property def max(self): return max(self.deque) @property def value(self): return self.deque[-1] def __str__(self): return self.fmt.format( median=self.median, avg=self.avg, global_avg=self.global_avg, max=self.max, value=self.value) def all_gather(data): """ Run all_gather on arbitrary picklable data (not necessarily tensors) Args: data: any picklable object Returns: list[data]: list of data gathered from each rank """ world_size = get_world_size() if world_size == 1: return [data] # serialized to a Tensor buffer = pickle.dumps(data) storage = torch.ByteStorage.from_buffer(buffer) tensor = torch.ByteTensor(storage).to("cuda") # obtain Tensor size of each rank local_size = torch.tensor([tensor.numel()], device="cuda") size_list = [torch.tensor([0], device="cuda") for _ in range(world_size)] dist.all_gather(size_list, local_size) size_list = [int(size.item()) for size in size_list] max_size = max(size_list) # receiving Tensor from all ranks # we pad the tensor because torch all_gather does not support # gathering tensors of different shapes tensor_list = [] for _ in size_list: tensor_list.append(torch.empty((max_size,), dtype=torch.uint8, device="cuda")) if local_size != max_size: padding = torch.empty(size=(max_size - local_size,), dtype=torch.uint8, device="cuda") tensor = torch.cat((tensor, padding), dim=0) dist.all_gather(tensor_list, tensor) data_list = [] for size, tensor in zip(size_list, tensor_list): buffer = tensor.cpu().numpy().tobytes()[:size] data_list.append(pickle.loads(buffer)) return data_list def reduce_dict(input_dict, average=True): """ Args: input_dict (dict): all the values will be reduced average (bool): whether to do average or sum Reduce the values in the dictionary from all processes so that all processes have the averaged results. Returns a dict with the same fields as input_dict, after reduction. """ world_size = get_world_size() if world_size < 2: return input_dict with torch.no_grad(): names = [] values = [] # sort the keys so that they are consistent across processes for k in sorted(input_dict.keys()): names.append(k) values.append(input_dict[k]) values = torch.stack(values, dim=0) dist.all_reduce(values) if average: values /= world_size reduced_dict = {k: v for k, v in zip(names, values)} return reduced_dict class MetricLogger(object): def __init__(self, delimiter="\t"): self.meters = defaultdict(SmoothedValue) self.delimiter = delimiter def update(self, **kwargs): for k, v in kwargs.items(): if isinstance(v, torch.Tensor): v = v.item() assert isinstance(v, (float, int)) self.meters[k].update(v) def __getattr__(self, attr): if attr in self.meters: return self.meters[attr] if attr in self.__dict__: return self.__dict__[attr] raise AttributeError("'{}' object has no attribute '{}'".format( type(self).__name__, attr)) def __str__(self): loss_str = [] for name, meter in self.meters.items(): # print(name, str(meter)) # import ipdb;ipdb.set_trace() if meter.count > 0: loss_str.append( "{}: {}".format(name, str(meter)) ) return self.delimiter.join(loss_str) def synchronize_between_processes(self): for meter in self.meters.values(): meter.synchronize_between_processes() def add_meter(self, name, meter): self.meters[name] = meter def log_every(self, iterable, print_freq, header=None, logger=None): if logger is None: print_func = print else: print_func = logger.info i = 0 if not header: header = '' start_time = time.time() end = time.time() iter_time = SmoothedValue(fmt='{avg:.4f}') data_time = SmoothedValue(fmt='{avg:.4f}') space_fmt = ':' + str(len(str(len(iterable)))) + 'd' if torch.cuda.is_available(): log_msg = self.delimiter.join([ header, '[{0' + space_fmt + '}/{1}]', 'eta: {eta}', '{meters}', 'time: {time}', 'data: {data}', 'max mem: {memory:.0f}' ]) else: log_msg = self.delimiter.join([ header, '[{0' + space_fmt + '}/{1}]', 'eta: {eta}', '{meters}', 'time: {time}', 'data: {data}' ]) MB = 1024.0 * 1024.0 for obj in iterable: data_time.update(time.time() - end) yield obj iter_time.update(time.time() - end) if i % print_freq == 0 or i == len(iterable) - 1: eta_seconds = iter_time.global_avg * (len(iterable) - i) eta_string = str(datetime.timedelta(seconds=int(eta_seconds))) if torch.cuda.is_available(): print_func(log_msg.format( i, len(iterable), eta=eta_string, meters=str(self), time=str(iter_time), data=str(data_time), memory=torch.cuda.max_memory_allocated() / MB)) else: print_func(log_msg.format( i, len(iterable), eta=eta_string, meters=str(self), time=str(iter_time), data=str(data_time))) i += 1 end = time.time() total_time = time.time() - start_time total_time_str = str(datetime.timedelta(seconds=int(total_time))) print_func('{} Total time: {} ({:.4f} s / it)'.format( header, total_time_str, total_time / len(iterable))) def get_sha(): cwd = os.path.dirname(os.path.abspath(__file__)) def _run(command): return subprocess.check_output(command, cwd=cwd).decode('ascii').strip() sha = 'N/A' diff = "clean" branch = 'N/A' try: sha = _run(['git', 'rev-parse', 'HEAD']) subprocess.check_output(['git', 'diff'], cwd=cwd) diff = _run(['git', 'diff-index', 'HEAD']) diff = "has uncommited changes" if diff else "clean" branch = _run(['git', 'rev-parse', '--abbrev-ref', 'HEAD']) except Exception: pass message = f"sha: {sha}, status: {diff}, branch: {branch}" return message def setup_for_distributed(is_master): """ This function disables printing when not in master process """ import builtins as __builtin__ builtin_print = __builtin__.print def print(*args, **kwargs): force = kwargs.pop('force', False) if is_master or force: builtin_print(*args, **kwargs) __builtin__.print = print def is_dist_avail_and_initialized(): if not dist.is_available(): return False if not dist.is_initialized(): return False return True def get_world_size(): if not is_dist_avail_and_initialized(): return 1 return dist.get_world_size() def get_rank(): if not is_dist_avail_and_initialized(): return 0 return dist.get_rank() def is_main_process(): return get_rank() == 0 def save_on_master(*args, **kwargs): if is_main_process(): torch.save(*args, **kwargs) def init_distributed_mode(args): if 'WORLD_SIZE' in os.environ and os.environ['WORLD_SIZE'] != '': # 'RANK' in os.environ and # args.rank = int(os.environ["RANK"]) # args.world_size = int(os.environ['WORLD_SIZE']) # args.gpu = args.local_rank = int(os.environ['LOCAL_RANK']) # launch by torch.distributed.launch # Single node # python -m torch.distributed.launch --nproc_per_node=8 main.py --world-size 1 --rank 0 ... # Multi nodes # python -m torch.distributed.launch --nproc_per_node=8 main.py --world-size 2 --rank 0 --dist-url 'tcp://IP_OF_NODE0:FREEPORT' ... # python -m torch.distributed.launch --nproc_per_node=8 main.py --world-size 2 --rank 1 --dist-url 'tcp://IP_OF_NODE0:FREEPORT' ... local_world_size = int(os.environ['WORLD_SIZE']) args.world_size = args.world_size * local_world_size args.gpu = args.local_rank = int(os.environ['LOCAL_RANK']) args.rank = args.rank * local_world_size + args.local_rank print('world size: {}, rank: {}, local rank: {}'.format(args.world_size, args.rank, args.local_rank)) print(json.dumps(dict(os.environ), indent=2)) elif 'SLURM_PROCID' in os.environ: args.rank = int(os.environ['SLURM_PROCID']) args.gpu = args.local_rank = int(os.environ['SLURM_LOCALID']) args.world_size = int(os.environ['SLURM_NPROCS']) print('world size: {}, world rank: {}, local rank: {}, device_count: {}'.format(args.world_size, args.rank, args.local_rank, torch.cuda.device_count())) else: print('Not using distributed mode') args.distributed = False args.world_size = 1 args.rank = 0 args.local_rank = 0 return print("world_size:{} rank:{} local_rank:{}".format(args.world_size, args.rank, args.local_rank)) args.distributed = True torch.cuda.set_device(args.local_rank) args.dist_backend = 'nccl' print('| distributed init (rank {}): {}'.format(args.rank, args.dist_url), flush=True) torch.distributed.init_process_group(backend=args.dist_backend, init_method=args.dist_url, world_size=args.world_size, rank=args.rank) print("Before torch.distributed.barrier()") torch.distributed.barrier() print("End torch.distributed.barrier()") setup_for_distributed(args.rank == 0) def masks_to_boxes(masks): """Compute the bounding boxes around the provided masks The masks should be in format [N, H, W] where N is the number of masks, (H, W) are the spatial dimensions. Returns a [N, 4] tensors, with the boxes in xyxy format """ if masks.numel() == 0: return torch.zeros((0, 4), device=masks.device) h, w = masks.shape[-2:] y = torch.arange(0, h, dtype=torch.float) x = torch.arange(0, w, dtype=torch.float) y, x = torch.meshgrid(y, x) y = y.to(masks) x = x.to(masks) x_mask = ((masks>128) * x.unsqueeze(0)) x_max = x_mask.flatten(1).max(-1)[0] x_min = x_mask.masked_fill(~(masks>128), 1e8).flatten(1).min(-1)[0] y_mask = ((masks>128) * y.unsqueeze(0)) y_max = y_mask.flatten(1).max(-1)[0] y_min = y_mask.masked_fill(~(masks>128), 1e8).flatten(1).min(-1)[0] return torch.stack([x_min, y_min, x_max, y_max], 1) def box_cxcywh_to_xyxy(x): x_c, y_c, w, h = x.unbind(-1) b = [(x_c - 0.5 * w), (y_c - 0.5 * h), (x_c + 0.5 * w), (y_c + 0.5 * h)] return torch.stack(b, dim=-1) def box_xyxy_to_cxcywh(x): x0, y0, x1, y1 = x.unbind(-1) b = [(x0 + x1) / 2, (y0 + y1) / 2, (x1 - x0), (y1 - y0)] return torch.stack(b, dim=-1) def box_noise(boxes, box_noise_scale=0): known_bbox_expand = box_xyxy_to_cxcywh(boxes) diff = torch.zeros_like(known_bbox_expand) diff[:, :2] = known_bbox_expand[:, 2:] / 2 diff[:, 2:] = known_bbox_expand[:, 2:] known_bbox_expand += torch.mul((torch.rand_like(known_bbox_expand) * 2 - 1.0),diff).cuda() * box_noise_scale boxes = box_cxcywh_to_xyxy(known_bbox_expand) boxes = boxes.clamp(min=0.0, max=1024) return boxes def masks_sample_points(masks,k=10): """Sample points on mask """ if masks.numel() == 0: return torch.zeros((0, 2), device=masks.device) h, w = masks.shape[-2:] y = torch.arange(0, h, dtype=torch.float) x = torch.arange(0, w, dtype=torch.float) y, x = torch.meshgrid(y, x) y = y.to(masks) x = x.to(masks) # k = 10 samples = [] for b_i in range(len(masks)): select_mask = (masks[b_i]>128) x_idx = torch.masked_select(x,select_mask) y_idx = torch.masked_select(y,select_mask) perm = torch.randperm(x_idx.size(0)) idx = perm[:k] samples_x = x_idx[idx] samples_y = y_idx[idx] samples_xy = torch.cat((samples_x[:,None],samples_y[:,None]),dim=1) samples.append(samples_xy) samples = torch.stack(samples) return samples # Add noise to mask input # From Mask Transfiner https://github.com/SysCV/transfiner def masks_noise(masks): def get_incoherent_mask(input_masks, sfact): mask = input_masks.float() w = input_masks.shape[-1] h = input_masks.shape[-2] mask_small = F.interpolate(mask, (h//sfact, w//sfact), mode='bilinear') mask_recover = F.interpolate(mask_small, (h, w), mode='bilinear') mask_residue = (mask - mask_recover).abs() mask_residue = (mask_residue >= 0.01).float() return mask_residue gt_masks_vector = masks / 255 mask_noise = torch.randn(gt_masks_vector.shape, device= gt_masks_vector.device) * 1.0 inc_masks = get_incoherent_mask(gt_masks_vector, 8) gt_masks_vector = ((gt_masks_vector + mask_noise * inc_masks) > 0.5).float() gt_masks_vector = gt_masks_vector * 255 return gt_masks_vector def mask_iou(pred_label,label): ''' calculate mask iou for pred_label and gt_label ''' pred_label = (pred_label>0.5)[0].int() label = (label>0.5)[0].int() intersection = ((label * pred_label) > 0).sum() union = ((label + pred_label) > 0).sum() return intersection / union # General util function to get the boundary of a binary mask. # https://gist.github.com/bowenc0221/71f7a02afee92646ca05efeeb14d687d def mask_to_boundary(mask, dilation_ratio=0.02): """ Convert binary mask to boundary mask. :param mask (numpy array, uint8): binary mask :param dilation_ratio (float): ratio to calculate dilation = dilation_ratio * image_diagonal :return: boundary mask (numpy array) """ h, w = mask.shape img_diag = np.sqrt(h ** 2 + w ** 2) dilation = int(round(dilation_ratio * img_diag)) if dilation < 1: dilation = 1 # Pad image so mask truncated by the image border is also considered as boundary. new_mask = cv2.copyMakeBorder(mask, 1, 1, 1, 1, cv2.BORDER_CONSTANT, value=0) kernel = np.ones((3, 3), dtype=np.uint8) new_mask_erode = cv2.erode(new_mask, kernel, iterations=dilation) mask_erode = new_mask_erode[1 : h + 1, 1 : w + 1] # G_d intersects G in the paper. return mask - mask_erode def boundary_iou(gt, dt, dilation_ratio=0.02): """ Compute boundary iou between two binary masks. :param gt (numpy array, uint8): binary mask :param dt (numpy array, uint8): binary mask :param dilation_ratio (float): ratio to calculate dilation = dilation_ratio * image_diagonal :return: boundary iou (float) """ device = gt.device dt = (dt>0.5)[0].cpu().byte().numpy() gt = (gt>0.5)[0].cpu().byte().numpy() gt_boundary = mask_to_boundary(gt, dilation_ratio) dt_boundary = mask_to_boundary(dt, dilation_ratio) intersection = ((gt_boundary * dt_boundary) > 0).sum() union = ((gt_boundary + dt_boundary) > 0).sum() boundary_iou = intersection / union return torch.tensor(boundary_iou).float().to(device) def bbox_overlaps(bboxes1, bboxes2, mode='iou', is_aligned=False, eps=1e-6): """Calculate overlap between two set of bboxes. If ``is_aligned`` is ``False``, then calculate the ious between each bbox of bboxes1 and bboxes2, otherwise the ious between each aligned pair of bboxes1 and bboxes2. Args: bboxes1 (Tensor): shape (m, 4) in format or empty. bboxes2 (Tensor): shape (n, 4) in format or empty. If is_aligned is ``True``, then m and n must be equal. mode (str): "iou" (intersection over union) or iof (intersection over foreground). Returns: ious(Tensor): shape (m, n) if is_aligned == False else shape (m, 1) Example: >>> bboxes1 = torch.FloatTensor([ >>> [0, 0, 10, 10], >>> [10, 10, 20, 20], >>> [32, 32, 38, 42], >>> ]) >>> bboxes2 = torch.FloatTensor([ >>> [0, 0, 10, 20], >>> [0, 10, 10, 19], >>> [10, 10, 20, 20], >>> ]) >>> bbox_overlaps(bboxes1, bboxes2) tensor([[0.5000, 0.0000, 0.0000], [0.0000, 0.0000, 1.0000], [0.0000, 0.0000, 0.0000]]) Example: >>> empty = torch.FloatTensor([]) >>> nonempty = torch.FloatTensor([ >>> [0, 0, 10, 9], >>> ]) >>> assert tuple(bbox_overlaps(empty, nonempty).shape) == (0, 1) >>> assert tuple(bbox_overlaps(nonempty, empty).shape) == (1, 0) >>> assert tuple(bbox_overlaps(empty, empty).shape) == (0, 0) """ assert mode in ['iou', 'iof'] # Either the boxes are empty or the length of boxes's last dimension is 4 assert (bboxes1.size(-1) == 4 or bboxes1.size(0) == 0) assert (bboxes2.size(-1) == 4 or bboxes2.size(0) == 0) rows = bboxes1.size(0) cols = bboxes2.size(0) if is_aligned: assert rows == cols if rows * cols == 0: return bboxes1.new(rows, 1) if is_aligned else bboxes1.new(rows, cols) if is_aligned: lt = torch.max(bboxes1[:, :2], bboxes2[:, :2]) # [rows, 2] rb = torch.min(bboxes1[:, 2:], bboxes2[:, 2:]) # [rows, 2] wh = (rb - lt).clamp(min=0) # [rows, 2] overlap = wh[:, 0] * wh[:, 1] area1 = (bboxes1[:, 2] - bboxes1[:, 0]) * (bboxes1[:, 3] - bboxes1[:, 1]) if mode == 'iou': area2 = (bboxes2[:, 2] - bboxes2[:, 0]) * (bboxes2[:, 3] - bboxes2[:, 1]) union = area1 + area2 - overlap else: union = area1 else: lt = torch.max(bboxes1[:, None, :2], bboxes2[:, :2]) # [rows, cols, 2] rb = torch.min(bboxes1[:, None, 2:], bboxes2[:, 2:]) # [rows, cols, 2] wh = (rb - lt).clamp(min=0) # [rows, cols, 2] overlap = wh[:, :, 0] * wh[:, :, 1] area1 = (bboxes1[:, 2] - bboxes1[:, 0]) * (bboxes1[:, 3] - bboxes1[:, 1]) if mode == 'iou': area2 = (bboxes2[:, 2] - bboxes2[:, 0]) * (bboxes2[:, 3] - bboxes2[:, 1]) union = area1[:, None] + area2 - overlap else: union = area1[:, None] eps = union.new_tensor([eps]) union = torch.max(union, eps) ious = overlap / union return ious def bbox_oiou(target, pred, eps=1e-7): # overlap lt = torch.max(pred[:, :2], target[:, :2]) rb = torch.min(pred[:, 2:], target[:, 2:]) wh = (rb - lt).clamp(min=0) overlap = wh[:, 0] * wh[:, 1] # union ap = (target[:, 2] - target[:, 0]) * (target[:, 3] - target[:, 1]) # IoU ious = overlap / ap return ious