commit EVF-SAMUltra node
This commit is contained in:
@@ -102,6 +102,8 @@ When this error has occurred, please check the network environment.
|
||||
## Update
|
||||
<font size="4">**If the dependency package error after updating, please reinstall the relevant dependency packages. </font><br />
|
||||
|
||||
* Commit [EVF-SAMUltra](#EVFSAMUltra) node, it is implementation of [EVF-SAM](https://github.com/hustvl/EVF-SAM) in ComfyUI. Please download model files from [BaiduNetdisk](https://pan.baidu.com/s/1EvaxgKcCxUpMbYKzLnEx9w?pwd=69bn) or [huggingface/EVF-SAM2](https://huggingface.co/YxZhang/evf-sam2/tree/main), [huggingface/EVF-SAM](https://huggingface.co/YxZhang/evf-sam/tree/main) to ```ComfyUI/models/EVF-SAM``` folder(save the models in their respective subdirectories).
|
||||
Due to the introduction of new dependencies package, after the plugin upgrade, please execute ```install_dequirements.bat``` to install.
|
||||
* Commit [ImageTaggerSave](#ImageTaggerSave) and [ImageAutoCropV3](#ImageAutoCropV3) nodes. Used to implement the automatic trimming and marking workflow for the training set (the workflow ```image_tagger_save.json``` is located in the workflow directory).
|
||||
* Commit [CheckMaskV2](#CheckMaskV2) node, Added the ```simple``` method to detect masks more quickly.
|
||||
* Commit [ImageReel](#ImageReel) and [ImageReelComposite](#ImageReelComposite) nodes to composite multiple images on a canvas.
|
||||
@@ -1510,7 +1512,7 @@ Output:
|
||||
* y: The y-coordinate of the top left corner position.
|
||||
|
||||
|
||||
## <a id="table1">Ultra</a>节点组
|
||||
## <a id="table1">Ultra</a> Nodes
|
||||

|
||||
Nodes that use ultra fine edge masking processing methods, the latest version of nodes includes: SegmentAnythingUltraV2, RmBgUltraV2, BiRefNetUltra, PersonMaskUltraV2, SegformerB2ClothesUltra and MaskEdgeUltraDetailV2.
|
||||
There are three edge processing methods for these nodes:
|
||||
@@ -1566,6 +1568,28 @@ On the basis of SegmentAnythingUltra, the following changes have been made:
|
||||
* device: Set whether the VitMatte to use cuda.
|
||||
* max_megapixels: Set the maximum size for VitMate operations.
|
||||
|
||||
### <a id="table1">EVF-SAMUltra</a>
|
||||
This node is implementation of [EVF-SAM](https://github.com/hustvl/EVF-SAM) in ComfyUI.
|
||||
*Please download model files from [BaiduNetdisk](https://pan.baidu.com/s/1EvaxgKcCxUpMbYKzLnEx9w?pwd=69bn) or [huggingface/EVF-SAM2](https://huggingface.co/YxZhang/evf-sam2/tree/main), [huggingface/EVF-SAM](https://huggingface.co/YxZhang/evf-sam/tree/main) to ```ComfyUI/models/EVF-SAM``` folder(save the models in their respective subdirectories).
|
||||

|
||||
|
||||
Node Options:
|
||||

|
||||
|
||||
* image: The input image.
|
||||
* model: Select the model. Currently, there are options for evf-sam2 and evf sam.
|
||||
* presicion: Model accuracy can be selected from fp16, bf16, and fp32.
|
||||
* load_in_bit: Load the model with positional accuracy. You can choose from 16, 8, and 4.
|
||||
* pormpt: Prompt words used for segmentation.
|
||||
* detail_method: Edge processing methods. provides VITMatte, VITMatte(local), PyMatting, GuidedFilter. If the model has been downloaded after the first use of VITMatte, you can use VITMatte (local) afterwards.
|
||||
* detail_erode: Mask the erosion range inward from the edge. the larger the value, the larger the range of inward repair.
|
||||
* detail_dilate: The edge of the mask expands outward. the larger the value, the wider the range of outward repair.
|
||||
* black_point: Edge black sampling threshold.
|
||||
* white_point: Edge white sampling threshold.
|
||||
* process_detail: Set to false here will skip edge processing to save runtime.
|
||||
* device: Set whether the VitMatte to use cuda.
|
||||
* max_megapixels: Set the maximum size for VitMate operations.
|
||||
|
||||
### <a id="table1">Florence2Ultra</a>
|
||||
Using the segmentation function of the Florence2 model, while also having ultra-high edge details.
|
||||
The code for this node section is from [spacepxl/ComfyUI-Florence-2](https://github.com/spacepxl/ComfyUI-Florence-2), thanks to the original author.
|
||||
|
||||
@@ -103,6 +103,8 @@ git clone https://github.com/chflame163/ComfyUI_LayerStyle.git
|
||||
## 更新说明
|
||||
<font size="4">**如果本插件更新后出现依赖包错误,请重新安装相关依赖包。
|
||||
|
||||
* 添加 [EVF-SAMUltra](#EVFSAMUltra) 节点,是[EVF-SAM](https://github.com/hustvl/EVF-SAM)在ComfyUI中的实现。请从[百度网盘](https://pan.baidu.com/s/1EvaxgKcCxUpMbYKzLnEx9w?pwd=69bn) 或者 [huggingface/EVF-SAM2](https://huggingface.co/YxZhang/evf-sam2/tree/main), [huggingface/EVF-SAM](https://huggingface.co/YxZhang/evf-sam/tree/main) 下载全部模型文件并复制到```ComfyUI/models/EVF-SAM```文件夹(请将模型保存在各自子目录中)。
|
||||
由于引入了新的依赖,插件升级后请执行```install_requirements.bat```安装新的依赖包。
|
||||
* 添加 [ImageTaggerSave](#ImageTaggerSave) 和 [ImageAutoCropV3](#ImageAutoCropV3) 节点,用于实现训练集自动裁切打标工作流(工作流```image_tagger_save_example.json```在workflow目录中)。
|
||||
* 添加 [CheckMaskV2](#CheckMaskV2) 节点,增加了```simple```方法以更快速检测遮罩。
|
||||
* 添加 [ImageReel ](#ImageReel) 和 [ImageReelComposit](#ImageReelComposit) 节点,可将多张图片显示在一起。
|
||||
@@ -1543,6 +1545,30 @@ SegmentAnythingUltra的V2升级版,增加了VITMatte边缘处理方法。
|
||||
* device: 设置是否使用cuda。
|
||||
* max_megapixels: 设置vitmatte运算的最大尺寸。
|
||||
|
||||
### <a id="table1">EVF-SAMUltra</a>
|
||||
本节点是[EVF-SAM](https://github.com/hustvl/EVF-SAM)在ComfyUI中的实现。
|
||||
*请从[百度网盘](https://pan.baidu.com/s/1EvaxgKcCxUpMbYKzLnEx9w?pwd=69bn) 或者 [huggingface/EVF-SAM2](https://huggingface.co/YxZhang/evf-sam2/tree/main), [huggingface/EVF-SAM](https://huggingface.co/YxZhang/evf-sam/tree/main) 下载全部模型文件并复制到```ComfyUI/models/EVF-SAM```文件夹(请将模型保存在各自子目录中)。
|
||||
|
||||

|
||||
|
||||
节点选项说明:
|
||||

|
||||
|
||||
* image: 图片输入。
|
||||
* model: 选择模型。目前有 evf-sam2 和 evf-sam 可选。
|
||||
* presicion: 模型精度,可选择fp16, bf16 和 fp32。
|
||||
* load_in_bit: 按位精度加载模型。可选择16,8 和 4。
|
||||
* pormpt: 用于分割的提示词。
|
||||
* detail_method: 边缘处理方法。提供了VITMatte, VITMatte(local), PyMatting, GuidedFilter。如果首次使用VITMatte后模型已经下载,之后可以使用VITMatte(local)。
|
||||
* detail_erode: 遮罩边缘向内侵蚀范围。数值越大,向内修复的范围越大。
|
||||
* detail_dilate: 遮罩边缘向外扩张范围。数值越大,向外修复的范围越大。
|
||||
* black_point: 边缘黑色采样阈值。
|
||||
* white_point: 边缘黑色采样阈值。
|
||||
* process_detail: 此处设为False将跳过边缘处理以节省运行时间。
|
||||
* device: 设置是否使用cuda。
|
||||
* max_megapixels: 设置vitmatte运算的最大尺寸。
|
||||
|
||||
|
||||
### <a id="table1">Florence2Ultra</a>
|
||||
使用 Florence2 模型的分割功能,同时具有超高的边缘细节。
|
||||
本节点部分的代码来自[spacepxl/ComfyUI-Florence-2](https://github.com/spacepxl/ComfyUI-Florence-2),感谢原作者。
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 436 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 125 KiB |
@@ -2,15 +2,21 @@
|
||||
|
||||
set "python_exec=..\..\..\python_embeded\python.exe"
|
||||
set "repair_dependency_txt=%~dp0\repair_dependency_list.txt"
|
||||
set "requirements_txt=%~dp0\requirements.txt"
|
||||
|
||||
echo Installing with ComfyUI Portable
|
||||
echo .
|
||||
echo Install whl...
|
||||
%python_exec% -s -m pip install ./whl/docopt-0.6.2-py2.py3-none-any.whl
|
||||
%python_exec% -s -m pip install ./whl/hydra_core-1.3.2-py3-none-any.whl
|
||||
|
||||
echo .
|
||||
echo Install requirement.txt...
|
||||
%python_exec% -s -m pip install -r ./requirements.txt
|
||||
|
||||
for /f "delims=" %%i in (%requirements_txt%) do (
|
||||
%python_exec% -s -m pip install "%%i"
|
||||
)
|
||||
|
||||
|
||||
echo .
|
||||
echo Fixing Dependency Package...
|
||||
|
||||
@@ -2,15 +2,19 @@
|
||||
|
||||
set "python_exec=..\..\python\python.exe"
|
||||
set "repair_dependency_txt=%~dp0\repair_dependency_list.txt"
|
||||
set "requirements_txt=%~dp0\requirements.txt"
|
||||
|
||||
echo Installing with ComfyUI Portable
|
||||
echo .
|
||||
echo Install whl...
|
||||
%python_exec% -s -m pip install ./whl/docopt-0.6.2-py2.py3-none-any.whl
|
||||
%python_exec% -s -m pip install ./whl/hydra_core-1.3.2-py3-none-any.whl
|
||||
|
||||
echo .
|
||||
echo Install requirement.txt...
|
||||
%python_exec% -s -m pip install -r ./requirements.txt
|
||||
for /f "delims=" %%i in (%requirements_txt%) do (
|
||||
%python_exec% -s -m pip install "%%i"
|
||||
)
|
||||
|
||||
echo .
|
||||
echo Fixing Dependency Package...
|
||||
|
||||
@@ -0,0 +1,145 @@
|
||||
import os
|
||||
import sys
|
||||
from PIL import Image
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torchvision import transforms
|
||||
from torchvision.transforms.functional import InterpolationMode
|
||||
from transformers import AutoTokenizer, BitsAndBytesConfig
|
||||
from model.segment_anything.utils.transforms import ResizeLongestSide
|
||||
|
||||
def sam_preprocess(
|
||||
x: np.ndarray,
|
||||
pixel_mean=torch.Tensor([123.675, 116.28, 103.53]).view(-1, 1, 1),
|
||||
pixel_std=torch.Tensor([58.395, 57.12, 57.375]).view(-1, 1, 1),
|
||||
img_size=1024,
|
||||
model_type="ori") -> torch.Tensor:
|
||||
'''
|
||||
preprocess of Segment Anything Model, including scaling, normalization and padding.
|
||||
preprocess differs between SAM and Effi-SAM, where Effi-SAM use no padding.
|
||||
input: ndarray
|
||||
output: torch.Tensor
|
||||
'''
|
||||
assert img_size==1024, \
|
||||
"both SAM and Effi-SAM receive images of size 1024^2, don't change this setting unless you're sure that your employed model works well with another size."
|
||||
x = ResizeLongestSide(img_size).apply_image(x)
|
||||
resize_shape = x.shape[:2]
|
||||
x = torch.from_numpy(x).permute(2,0,1).contiguous()
|
||||
|
||||
# Normalize colors
|
||||
x = (x - pixel_mean) / pixel_std
|
||||
if model_type=="effi" or model_type=="sam2":
|
||||
x = F.interpolate(x.unsqueeze(0), (img_size, img_size), mode="bilinear").squeeze(0)
|
||||
else:
|
||||
# Pad
|
||||
h, w = x.shape[-2:]
|
||||
padh = img_size - h
|
||||
padw = img_size - w
|
||||
x = F.pad(x, (0, padw, 0, padh))
|
||||
return x, resize_shape
|
||||
|
||||
def beit3_preprocess(x: np.ndarray, img_size=224) -> torch.Tensor:
|
||||
'''
|
||||
preprocess for BEIT-3 model.
|
||||
input: ndarray
|
||||
output: torch.Tensor
|
||||
'''
|
||||
beit_preprocess = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
transforms.Resize((img_size, img_size), interpolation=InterpolationMode.BICUBIC),
|
||||
transforms.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5))
|
||||
])
|
||||
return beit_preprocess(x)
|
||||
|
||||
def init_models(model_path:str, model_type:str, precision:str, load_in_bit:int=16):
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
model_path,
|
||||
padding_side="right",
|
||||
use_fast=False,
|
||||
)
|
||||
|
||||
torch_dtype = torch.float32
|
||||
if precision == "bf16":
|
||||
torch_dtype = torch.bfloat16
|
||||
elif precision == "fp16":
|
||||
torch_dtype = torch.half
|
||||
|
||||
kwargs = {"torch_dtype": torch_dtype}
|
||||
|
||||
if load_in_bit==4:
|
||||
kwargs.update(
|
||||
{
|
||||
"torch_dtype": torch.half,
|
||||
"quantization_config": BitsAndBytesConfig(
|
||||
llm_int8_skip_modules=["visual_model"],
|
||||
load_in_4bit=True,
|
||||
bnb_4bit_compute_dtype=torch.float16,
|
||||
bnb_4bit_use_double_quant=True,
|
||||
bnb_4bit_quant_type="nf4",
|
||||
),
|
||||
}
|
||||
)
|
||||
elif load_in_bit==8:
|
||||
kwargs.update(
|
||||
{
|
||||
"torch_dtype": torch.half,
|
||||
"quantization_config": BitsAndBytesConfig(
|
||||
llm_int8_skip_modules=["visual_model"],
|
||||
load_in_8bit=True,
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
if model_type=="ori":
|
||||
from model.evf_sam import EvfSamModel
|
||||
model = EvfSamModel.from_pretrained(
|
||||
model_path, low_cpu_mem_usage=True, **kwargs
|
||||
)
|
||||
elif model_type=="effi":
|
||||
from model.evf_effisam import EvfEffiSamModel
|
||||
model = EvfEffiSamModel.from_pretrained(
|
||||
model_path, low_cpu_mem_usage=True, **kwargs
|
||||
)
|
||||
elif model_type=="sam2":
|
||||
from model.evf_sam2 import EvfSam2Model
|
||||
model = EvfSam2Model.from_pretrained(
|
||||
model_path, low_cpu_mem_usage=True, **kwargs
|
||||
)
|
||||
|
||||
if load_in_bit > 8 and torch.cuda.is_available():
|
||||
model = model.cuda()
|
||||
model.eval()
|
||||
|
||||
return tokenizer, model
|
||||
|
||||
def evf_sam_main(model_path:str, model_type:str, precision:str, load_in_bit:int, image:Image, prompt:str, ):
|
||||
|
||||
image_size = 224
|
||||
# initialize model and tokenizer
|
||||
tokenizer, model = init_models(model_path, model_type, precision, load_in_bit)
|
||||
|
||||
# preprocess
|
||||
image_np = np.asarray(image)
|
||||
image_np = cv2.cvtColor(image_np, cv2.COLOR_BGR2RGB)
|
||||
original_size_list = [image_np.shape[:2]]
|
||||
image_beit = beit3_preprocess(image_np, image_size).to(dtype=model.dtype, device=model.device)
|
||||
image_sam, resize_shape = sam_preprocess(image_np, model_type=model_type)
|
||||
image_sam = image_sam.to(dtype=model.dtype, device=model.device)
|
||||
input_ids = tokenizer(prompt, return_tensors="pt")["input_ids"].to(device=model.device)
|
||||
|
||||
# infer
|
||||
pred_mask = model.inference(
|
||||
image_sam.unsqueeze(0),
|
||||
image_beit.unsqueeze(0),
|
||||
input_ids,
|
||||
resize_list=[resize_shape],
|
||||
original_size_list=original_size_list,
|
||||
)
|
||||
|
||||
pred_mask = pred_mask.detach().cpu().numpy()[0]
|
||||
pred_mask = (pred_mask > 0).astype(np.uint8) * 255
|
||||
out_put_image = Image.fromarray(pred_mask.squeeze(), mode="L")
|
||||
|
||||
return out_put_image
|
||||
@@ -0,0 +1,7 @@
|
||||
# Copyright (c) Meta Platforms, Inc. and affiliates.
|
||||
# All rights reserved.
|
||||
|
||||
from .build_efficient_sam import (
|
||||
build_efficient_sam_vitt,
|
||||
build_efficient_sam_vits,
|
||||
)
|
||||
@@ -0,0 +1,22 @@
|
||||
# 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.
|
||||
|
||||
from .efficient_sam import build_efficient_sam
|
||||
|
||||
def build_efficient_sam_vitt(checkpoint=None):
|
||||
return build_efficient_sam(
|
||||
encoder_patch_embed_dim=192,
|
||||
encoder_num_heads=3,
|
||||
checkpoint=checkpoint,
|
||||
).eval()
|
||||
|
||||
|
||||
def build_efficient_sam_vits(checkpoint=None):
|
||||
return build_efficient_sam(
|
||||
encoder_patch_embed_dim=384,
|
||||
encoder_num_heads=6,
|
||||
checkpoint=checkpoint,
|
||||
).eval()
|
||||
@@ -0,0 +1,306 @@
|
||||
# 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 math
|
||||
from typing import Any, List, Tuple, Type
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from torch import nn, Tensor
|
||||
|
||||
from .efficient_sam_decoder import MaskDecoder, PromptEncoder
|
||||
from .efficient_sam_encoder import ImageEncoderViT
|
||||
from .two_way_transformer import TwoWayAttentionBlock, TwoWayTransformer
|
||||
|
||||
class EfficientSam(nn.Module):
|
||||
mask_threshold: float = 0.0
|
||||
image_format: str = "RGB"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
image_encoder: ImageEncoderViT,
|
||||
prompt_encoder: PromptEncoder,
|
||||
decoder_max_num_input_points: int,
|
||||
mask_decoder: MaskDecoder,
|
||||
pixel_mean: List[float] = [0.485, 0.456, 0.406],
|
||||
pixel_std: List[float] = [0.229, 0.224, 0.225],
|
||||
) -> None:
|
||||
"""
|
||||
SAM predicts object masks from an image and input prompts.
|
||||
|
||||
Arguments:
|
||||
image_encoder (ImageEncoderViT): The backbone used to encode the
|
||||
image into image embeddings that allow for efficient mask prediction.
|
||||
prompt_encoder (PromptEncoder): Encodes various types of input prompts.
|
||||
mask_decoder (MaskDecoder): Predicts masks from the image embeddings
|
||||
and encoded prompts.
|
||||
pixel_mean (list(float)): Mean values for normalizing pixels in the input image.
|
||||
pixel_std (list(float)): Std values for normalizing pixels in the input image.
|
||||
"""
|
||||
super().__init__()
|
||||
self.image_encoder = image_encoder
|
||||
self.prompt_encoder = prompt_encoder
|
||||
self.decoder_max_num_input_points = decoder_max_num_input_points
|
||||
self.mask_decoder = mask_decoder
|
||||
self.register_buffer(
|
||||
"pixel_mean", torch.Tensor(pixel_mean).view(1, 3, 1, 1), False
|
||||
)
|
||||
self.register_buffer(
|
||||
"pixel_std", torch.Tensor(pixel_std).view(1, 3, 1, 1), False
|
||||
)
|
||||
|
||||
@torch.jit.export
|
||||
def predict_masks(
|
||||
self,
|
||||
image_embeddings: torch.Tensor,
|
||||
batched_points: torch.Tensor,
|
||||
batched_point_labels: torch.Tensor,
|
||||
multimask_output: bool,
|
||||
input_h: int,
|
||||
input_w: int,
|
||||
output_h: int = -1,
|
||||
output_w: int = -1,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Predicts masks given image embeddings and prompts. This only runs the decoder.
|
||||
|
||||
Arguments:
|
||||
image_embeddings: A tensor of shape [B, C, H, W] or [B*max_num_queries, C, H, W]
|
||||
batched_points: A tensor of shape [B, max_num_queries, num_pts, 2]
|
||||
batched_point_labels: A tensor of shape [B, max_num_queries, num_pts]
|
||||
Returns:
|
||||
A tuple of two tensors:
|
||||
low_res_mask: A tensor of shape [B, max_num_queries, 256, 256] of predicted masks
|
||||
iou_predictions: A tensor of shape [B, max_num_queries] of estimated IOU scores
|
||||
"""
|
||||
|
||||
batch_size, max_num_queries, num_pts, _ = batched_points.shape
|
||||
num_pts = batched_points.shape[2]
|
||||
rescaled_batched_points = self.get_rescaled_pts(batched_points, input_h, input_w)
|
||||
|
||||
if num_pts > self.decoder_max_num_input_points:
|
||||
rescaled_batched_points = rescaled_batched_points[
|
||||
:, :, : self.decoder_max_num_input_points, :
|
||||
]
|
||||
batched_point_labels = batched_point_labels[
|
||||
:, :, : self.decoder_max_num_input_points
|
||||
]
|
||||
elif num_pts < self.decoder_max_num_input_points:
|
||||
rescaled_batched_points = F.pad(
|
||||
rescaled_batched_points,
|
||||
(0, 0, 0, self.decoder_max_num_input_points - num_pts),
|
||||
value=-1.0,
|
||||
)
|
||||
batched_point_labels = F.pad(
|
||||
batched_point_labels,
|
||||
(0, self.decoder_max_num_input_points - num_pts),
|
||||
value=-1.0,
|
||||
)
|
||||
|
||||
sparse_embeddings = self.prompt_encoder(
|
||||
rescaled_batched_points.reshape(
|
||||
batch_size * max_num_queries, self.decoder_max_num_input_points, 2
|
||||
),
|
||||
batched_point_labels.reshape(
|
||||
batch_size * max_num_queries, self.decoder_max_num_input_points
|
||||
),
|
||||
)
|
||||
|
||||
sparse_embeddings = sparse_embeddings.view(
|
||||
batch_size,
|
||||
max_num_queries,
|
||||
sparse_embeddings.shape[1],
|
||||
sparse_embeddings.shape[2],
|
||||
)
|
||||
low_res_masks, iou_predictions = self.mask_decoder(
|
||||
image_embeddings,
|
||||
self.prompt_encoder.get_dense_pe(),
|
||||
sparse_prompt_embeddings=sparse_embeddings,
|
||||
multimask_output=multimask_output,
|
||||
)
|
||||
_, num_predictions, low_res_size, _ = low_res_masks.shape
|
||||
|
||||
if output_w > 0 and output_h > 0:
|
||||
output_masks = F.interpolate(
|
||||
low_res_masks, (output_h, output_w), mode="bicubic"
|
||||
)
|
||||
output_masks = torch.reshape(
|
||||
output_masks,
|
||||
(batch_size, max_num_queries, num_predictions, output_h, output_w),
|
||||
)
|
||||
else:
|
||||
output_masks = torch.reshape(
|
||||
low_res_masks,
|
||||
(
|
||||
batch_size,
|
||||
max_num_queries,
|
||||
num_predictions,
|
||||
low_res_size,
|
||||
low_res_size,
|
||||
),
|
||||
)
|
||||
iou_predictions = torch.reshape(
|
||||
iou_predictions, (batch_size, max_num_queries, num_predictions)
|
||||
)
|
||||
return output_masks, iou_predictions
|
||||
|
||||
def get_rescaled_pts(self, batched_points: torch.Tensor, input_h: int, input_w: int):
|
||||
return torch.stack(
|
||||
[
|
||||
torch.where(
|
||||
batched_points[..., 0] >= 0,
|
||||
batched_points[..., 0] * self.image_encoder.img_size / input_w,
|
||||
-1.0,
|
||||
),
|
||||
torch.where(
|
||||
batched_points[..., 1] >= 0,
|
||||
batched_points[..., 1] * self.image_encoder.img_size / input_h,
|
||||
-1.0,
|
||||
),
|
||||
],
|
||||
dim=-1,
|
||||
)
|
||||
|
||||
@torch.jit.export
|
||||
def get_image_embeddings(self, batched_images) -> torch.Tensor:
|
||||
"""
|
||||
Predicts masks end-to-end from provided images and prompts.
|
||||
If prompts are not known in advance, using SamPredictor is
|
||||
recommended over calling the model directly.
|
||||
|
||||
Arguments:
|
||||
batched_images: A tensor of shape [B, 3, H, W]
|
||||
Returns:
|
||||
List of image embeddings each of of shape [B, C(i), H(i), W(i)].
|
||||
The last embedding corresponds to the final layer.
|
||||
"""
|
||||
batched_images = self.preprocess(batched_images)
|
||||
return self.image_encoder(batched_images)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batched_images: torch.Tensor,
|
||||
batched_points: torch.Tensor,
|
||||
batched_point_labels: torch.Tensor,
|
||||
scale_to_original_image_size: bool = True,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Predicts masks end-to-end from provided images and prompts.
|
||||
If prompts are not known in advance, using SamPredictor is
|
||||
recommended over calling the model directly.
|
||||
|
||||
Arguments:
|
||||
batched_images: A tensor of shape [B, 3, H, W]
|
||||
batched_points: A tensor of shape [B, num_queries, max_num_pts, 2]
|
||||
batched_point_labels: A tensor of shape [B, num_queries, max_num_pts]
|
||||
|
||||
Returns:
|
||||
A list tuples of two tensors where the ith element is by considering the first i+1 points.
|
||||
low_res_mask: A tensor of shape [B, 256, 256] of predicted masks
|
||||
iou_predictions: A tensor of shape [B, max_num_queries] of estimated IOU scores
|
||||
"""
|
||||
batch_size, _, input_h, input_w = batched_images.shape
|
||||
image_embeddings = self.get_image_embeddings(batched_images)
|
||||
return self.predict_masks(
|
||||
image_embeddings,
|
||||
batched_points,
|
||||
batched_point_labels,
|
||||
multimask_output=True,
|
||||
input_h=input_h,
|
||||
input_w=input_w,
|
||||
output_h=input_h if scale_to_original_image_size else -1,
|
||||
output_w=input_w if scale_to_original_image_size else -1,
|
||||
)
|
||||
|
||||
def preprocess(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Normalize pixel values and pad to a square input."""
|
||||
if (
|
||||
x.shape[2] != self.image_encoder.img_size
|
||||
or x.shape[3] != self.image_encoder.img_size
|
||||
):
|
||||
x = F.interpolate(
|
||||
x,
|
||||
(self.image_encoder.img_size, self.image_encoder.img_size),
|
||||
mode="bilinear",
|
||||
)
|
||||
return (x - self.pixel_mean) / self.pixel_std
|
||||
|
||||
|
||||
def build_efficient_sam(encoder_patch_embed_dim, encoder_num_heads, checkpoint=None):
|
||||
img_size = 1024
|
||||
encoder_patch_size = 16
|
||||
encoder_depth = 12
|
||||
encoder_mlp_ratio = 4.0
|
||||
encoder_neck_dims = [256, 256]
|
||||
decoder_max_num_input_points = 6
|
||||
decoder_transformer_depth = 2
|
||||
decoder_transformer_mlp_dim = 2048
|
||||
decoder_num_heads = 8
|
||||
decoder_upscaling_layer_dims = [64, 32]
|
||||
num_multimask_outputs = 3
|
||||
iou_head_depth = 3
|
||||
iou_head_hidden_dim = 256
|
||||
activation = "gelu"
|
||||
normalization_type = "layer_norm"
|
||||
normalize_before_activation = False
|
||||
|
||||
assert activation == "relu" or activation == "gelu"
|
||||
if activation == "relu":
|
||||
activation_fn = nn.ReLU
|
||||
else:
|
||||
activation_fn = nn.GELU
|
||||
|
||||
image_encoder = ImageEncoderViT(
|
||||
img_size=img_size,
|
||||
patch_size=encoder_patch_size,
|
||||
in_chans=3,
|
||||
patch_embed_dim=encoder_patch_embed_dim,
|
||||
normalization_type=normalization_type,
|
||||
depth=encoder_depth,
|
||||
num_heads=encoder_num_heads,
|
||||
mlp_ratio=encoder_mlp_ratio,
|
||||
neck_dims=encoder_neck_dims,
|
||||
act_layer=activation_fn,
|
||||
)
|
||||
|
||||
image_embedding_size = image_encoder.image_embedding_size
|
||||
encoder_transformer_output_dim = image_encoder.transformer_output_dim
|
||||
|
||||
sam = EfficientSam(
|
||||
image_encoder=image_encoder,
|
||||
prompt_encoder=PromptEncoder(
|
||||
embed_dim=encoder_transformer_output_dim,
|
||||
image_embedding_size=(image_embedding_size, image_embedding_size),
|
||||
input_image_size=(img_size, img_size),
|
||||
),
|
||||
decoder_max_num_input_points=decoder_max_num_input_points,
|
||||
mask_decoder=MaskDecoder(
|
||||
transformer_dim=encoder_transformer_output_dim,
|
||||
transformer=TwoWayTransformer(
|
||||
depth=decoder_transformer_depth,
|
||||
embedding_dim=encoder_transformer_output_dim,
|
||||
num_heads=decoder_num_heads,
|
||||
mlp_dim=decoder_transformer_mlp_dim,
|
||||
activation=activation_fn,
|
||||
normalize_before_activation=normalize_before_activation,
|
||||
),
|
||||
num_multimask_outputs=num_multimask_outputs,
|
||||
activation=activation_fn,
|
||||
normalization_type=normalization_type,
|
||||
normalize_before_activation=normalize_before_activation,
|
||||
iou_head_depth=iou_head_depth - 1,
|
||||
iou_head_hidden_dim=iou_head_hidden_dim,
|
||||
upscaling_layer_dims=decoder_upscaling_layer_dims,
|
||||
),
|
||||
pixel_mean=[0.485, 0.456, 0.406],
|
||||
pixel_std=[0.229, 0.224, 0.225],
|
||||
)
|
||||
if checkpoint is not None:
|
||||
with open(checkpoint, "rb") as f:
|
||||
state_dict = torch.load(f, map_location="cpu")
|
||||
sam.load_state_dict(state_dict["model"])
|
||||
return sam
|
||||
@@ -0,0 +1,318 @@
|
||||
# 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.
|
||||
|
||||
from typing import List, Tuple, Type
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from .mlp import MLPBlock
|
||||
|
||||
|
||||
class PromptEncoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
embed_dim: int,
|
||||
image_embedding_size: Tuple[int, int],
|
||||
input_image_size: Tuple[int, int],
|
||||
) -> None:
|
||||
"""
|
||||
Encodes prompts for input to SAM's mask decoder.
|
||||
|
||||
Arguments:
|
||||
embed_dim (int): The prompts' embedding dimension
|
||||
image_embedding_size (tuple(int, int)): The spatial size of the
|
||||
image embedding, as (H, W).
|
||||
input_image_size (int): The padded size of the image as input
|
||||
to the image encoder, as (H, W).
|
||||
"""
|
||||
super().__init__()
|
||||
self.embed_dim = embed_dim
|
||||
self.input_image_size = input_image_size
|
||||
self.image_embedding_size = image_embedding_size
|
||||
self.pe_layer = PositionEmbeddingRandom(embed_dim // 2)
|
||||
self.invalid_points = nn.Embedding(1, embed_dim)
|
||||
self.point_embeddings = nn.Embedding(1, embed_dim)
|
||||
self.bbox_top_left_embeddings = nn.Embedding(1, embed_dim)
|
||||
self.bbox_bottom_right_embeddings = nn.Embedding(1, embed_dim)
|
||||
|
||||
def get_dense_pe(self) -> torch.Tensor:
|
||||
"""
|
||||
Returns the positional encoding used to encode point prompts,
|
||||
applied to a dense set of points the shape of the image encoding.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Positional encoding with shape
|
||||
1x(embed_dim)x(embedding_h)x(embedding_w)
|
||||
"""
|
||||
return self.pe_layer(self.image_embedding_size).unsqueeze(0)
|
||||
|
||||
def _embed_points(
|
||||
self,
|
||||
points: torch.Tensor,
|
||||
labels: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Embeds point prompts."""
|
||||
|
||||
points = points + 0.5 # Shift to center of pixel
|
||||
point_embedding = self.pe_layer.forward_with_coords(
|
||||
points, self.input_image_size
|
||||
)
|
||||
invalid_label_ids = torch.eq(labels, -1)[:,:,None]
|
||||
point_label_ids = torch.eq(labels, 1)[:,:,None]
|
||||
topleft_label_ids = torch.eq(labels, 2)[:,:,None]
|
||||
bottomright_label_ids = torch.eq(labels, 3)[:,:,None]
|
||||
point_embedding = point_embedding + self.invalid_points.weight[:,None,:] * invalid_label_ids
|
||||
point_embedding = point_embedding + self.point_embeddings.weight[:,None,:] * point_label_ids
|
||||
point_embedding = point_embedding + self.bbox_top_left_embeddings.weight[:,None,:] * topleft_label_ids
|
||||
point_embedding = point_embedding + self.bbox_bottom_right_embeddings.weight[:,None,:] * bottomright_label_ids
|
||||
return point_embedding
|
||||
|
||||
def forward(
|
||||
self,
|
||||
coords,
|
||||
labels,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Embeds different types of prompts, returning both sparse and dense
|
||||
embeddings.
|
||||
|
||||
Arguments:
|
||||
points: A tensor of shape [B, 2]
|
||||
labels: An integer tensor of shape [B] where each element is 1,2 or 3.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: sparse embeddings for the points and boxes, with shape
|
||||
BxNx(embed_dim), where N is determined by the number of input points
|
||||
and boxes.
|
||||
"""
|
||||
return self._embed_points(coords, labels)
|
||||
|
||||
|
||||
class PositionEmbeddingRandom(nn.Module):
|
||||
"""
|
||||
Positional encoding using random spatial frequencies.
|
||||
"""
|
||||
|
||||
def __init__(self, num_pos_feats: int) -> None:
|
||||
super().__init__()
|
||||
self.register_buffer(
|
||||
"positional_encoding_gaussian_matrix", torch.randn((2, num_pos_feats))
|
||||
)
|
||||
|
||||
def _pe_encoding(self, coords: torch.Tensor) -> torch.Tensor:
|
||||
"""Positionally encode points that are normalized to [0,1]."""
|
||||
# assuming coords are in [0, 1]^2 square and have d_1 x ... x d_n x 2 shape
|
||||
coords = 2 * coords - 1
|
||||
coords = coords @ self.positional_encoding_gaussian_matrix
|
||||
coords = 2 * np.pi * coords
|
||||
# outputs d_1 x ... x d_n x C shape
|
||||
return torch.cat([torch.sin(coords), torch.cos(coords)], dim=-1)
|
||||
|
||||
def forward(self, size: Tuple[int, int]) -> torch.Tensor:
|
||||
"""Generate positional encoding for a grid of the specified size."""
|
||||
h, w = size
|
||||
device = self.positional_encoding_gaussian_matrix.device
|
||||
grid = torch.ones([h, w], device=device, dtype=self.positional_encoding_gaussian_matrix.dtype)
|
||||
y_embed = grid.cumsum(dim=0) - 0.5
|
||||
x_embed = grid.cumsum(dim=1) - 0.5
|
||||
y_embed = y_embed / h
|
||||
x_embed = x_embed / w
|
||||
|
||||
pe = self._pe_encoding(torch.stack([x_embed, y_embed], dim=-1))
|
||||
return pe.permute(2, 0, 1) # C x H x W
|
||||
|
||||
def forward_with_coords(
|
||||
self, coords_input: torch.Tensor, image_size: Tuple[int, int]
|
||||
) -> torch.Tensor:
|
||||
"""Positionally encode points that are not normalized to [0,1]."""
|
||||
coords = coords_input.clone()
|
||||
coords[:, :, 0] = coords[:, :, 0] / image_size[1]
|
||||
coords[:, :, 1] = coords[:, :, 1] / image_size[0]
|
||||
# remove to(float) here, don't know why original implementation add this
|
||||
return self._pe_encoding(coords) # B x N x C
|
||||
|
||||
|
||||
class MaskDecoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
transformer_dim: int,
|
||||
transformer: nn.Module,
|
||||
num_multimask_outputs: int,
|
||||
activation: Type[nn.Module],
|
||||
normalization_type: str,
|
||||
normalize_before_activation: bool,
|
||||
iou_head_depth: int,
|
||||
iou_head_hidden_dim: int,
|
||||
upscaling_layer_dims: List[int],
|
||||
) -> None:
|
||||
"""
|
||||
Predicts masks given an image and prompt embeddings, using a
|
||||
transformer architecture.
|
||||
|
||||
Arguments:
|
||||
transformer_dim (int): the channel dimension of the transformer
|
||||
transformer (nn.Module): the transformer used to predict masks
|
||||
num_multimask_outputs (int): the number of masks to predict
|
||||
when disambiguating masks
|
||||
activation (nn.Module): the type of activation to use when
|
||||
upscaling masks
|
||||
iou_head_depth (int): the depth of the MLP used to predict
|
||||
mask quality
|
||||
iou_head_hidden_dim (int): the hidden dimension of the MLP
|
||||
used to predict mask quality
|
||||
"""
|
||||
super().__init__()
|
||||
self.transformer_dim = transformer_dim
|
||||
self.transformer = transformer
|
||||
|
||||
self.num_multimask_outputs = num_multimask_outputs
|
||||
|
||||
self.iou_token = nn.Embedding(1, transformer_dim)
|
||||
if num_multimask_outputs > 1:
|
||||
self.num_mask_tokens = num_multimask_outputs + 1
|
||||
else:
|
||||
self.num_mask_tokens = 1
|
||||
self.mask_tokens = nn.Embedding(self.num_mask_tokens, transformer_dim)
|
||||
output_dim_after_upscaling = transformer_dim
|
||||
|
||||
self.final_output_upscaling_layers = nn.ModuleList([])
|
||||
for idx, layer_dims in enumerate(upscaling_layer_dims):
|
||||
self.final_output_upscaling_layers.append(
|
||||
nn.Sequential(
|
||||
nn.ConvTranspose2d(
|
||||
output_dim_after_upscaling,
|
||||
layer_dims,
|
||||
kernel_size=2,
|
||||
stride=2,
|
||||
),
|
||||
nn.GroupNorm(1, layer_dims)
|
||||
if idx < len(upscaling_layer_dims) - 1
|
||||
else nn.Identity(),
|
||||
activation(),
|
||||
)
|
||||
)
|
||||
output_dim_after_upscaling = layer_dims
|
||||
|
||||
self.output_hypernetworks_mlps = nn.ModuleList(
|
||||
[
|
||||
MLPBlock(
|
||||
input_dim=transformer_dim,
|
||||
hidden_dim=transformer_dim,
|
||||
output_dim=output_dim_after_upscaling,
|
||||
num_layers=2,
|
||||
act=activation,
|
||||
)
|
||||
for i in range(self.num_mask_tokens)
|
||||
]
|
||||
)
|
||||
|
||||
self.iou_prediction_head = MLPBlock(
|
||||
input_dim=transformer_dim,
|
||||
hidden_dim=iou_head_hidden_dim,
|
||||
output_dim=self.num_mask_tokens,
|
||||
num_layers=iou_head_depth,
|
||||
act=activation,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
image_embeddings: torch.Tensor,
|
||||
image_pe: torch.Tensor,
|
||||
sparse_prompt_embeddings: torch.Tensor,
|
||||
multimask_output: bool,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Predict masks given image and prompt embeddings.
|
||||
|
||||
Arguments:
|
||||
image_embeddings: A tensor of shape [B, C, H, W] or [B*max_num_queries, C, H, W]
|
||||
image_pe (torch.Tensor): positional encoding with the shape of image_embeddings (the batch dimension is broadcastable).
|
||||
sparse_prompt_embeddings (torch.Tensor): the embeddings of the points and boxes
|
||||
multimask_output (bool): Whether to return multiple masks or a single
|
||||
mask.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: batched predicted masks
|
||||
torch.Tensor: batched predictions of mask quality
|
||||
"""
|
||||
|
||||
(
|
||||
batch_size,
|
||||
max_num_queries,
|
||||
sparse_embed_dim_1,
|
||||
sparse_embed_dim_2,
|
||||
) = sparse_prompt_embeddings.shape
|
||||
|
||||
(
|
||||
_,
|
||||
image_embed_dim_c,
|
||||
image_embed_dim_h,
|
||||
image_embed_dim_w,
|
||||
) = image_embeddings.shape
|
||||
|
||||
# Tile the image embedding for all queries.
|
||||
image_embeddings_tiled = torch.tile(
|
||||
image_embeddings[:, None, :, :, :], [1, max_num_queries, 1, 1, 1]
|
||||
).view(
|
||||
batch_size * max_num_queries,
|
||||
image_embed_dim_c,
|
||||
image_embed_dim_h,
|
||||
image_embed_dim_w,
|
||||
)
|
||||
sparse_prompt_embeddings = sparse_prompt_embeddings.reshape(
|
||||
batch_size * max_num_queries, sparse_embed_dim_1, sparse_embed_dim_2
|
||||
)
|
||||
masks, iou_pred = self.predict_masks(
|
||||
image_embeddings=image_embeddings_tiled,
|
||||
image_pe=image_pe,
|
||||
sparse_prompt_embeddings=sparse_prompt_embeddings,
|
||||
)
|
||||
|
||||
if multimask_output and self.num_multimask_outputs > 1:
|
||||
return masks[:, 1:, :], iou_pred[:, 1:]
|
||||
else:
|
||||
return masks[:, :1, :], iou_pred[:, :1]
|
||||
|
||||
def predict_masks(
|
||||
self,
|
||||
image_embeddings: torch.Tensor,
|
||||
image_pe: torch.Tensor,
|
||||
sparse_prompt_embeddings: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Predicts masks. See 'forward' for more details."""
|
||||
# Concatenate output tokens
|
||||
output_tokens = torch.cat(
|
||||
[self.iou_token.weight, self.mask_tokens.weight], dim=0
|
||||
)
|
||||
output_tokens = output_tokens.unsqueeze(0).expand(
|
||||
sparse_prompt_embeddings.size(0), -1, -1
|
||||
)
|
||||
tokens = torch.cat((output_tokens, sparse_prompt_embeddings), dim=1)
|
||||
# Expand per-image data in batch direction to be per-mask
|
||||
pos_src = torch.repeat_interleave(image_pe, tokens.shape[0], dim=0)
|
||||
b, c, h, w = image_embeddings.shape
|
||||
hs, src = self.transformer(image_embeddings, pos_src, tokens)
|
||||
iou_token_out = hs[:, 0, :]
|
||||
mask_tokens_out = hs[:, 1 : (1 + self.num_mask_tokens), :]
|
||||
|
||||
# Upscale mask embeddings and predict masks using the mask tokens
|
||||
upscaled_embedding = src.transpose(1, 2).view(b, c, h, w)
|
||||
|
||||
for upscaling_layer in self.final_output_upscaling_layers:
|
||||
upscaled_embedding = upscaling_layer(upscaled_embedding)
|
||||
hyper_in_list: List[torch.Tensor] = []
|
||||
for i, output_hypernetworks_mlp in enumerate(self.output_hypernetworks_mlps):
|
||||
hyper_in_list.append(output_hypernetworks_mlp(mask_tokens_out[:, i, :]))
|
||||
hyper_in = torch.stack(hyper_in_list, dim=1)
|
||||
b, c, h, w = upscaled_embedding.shape
|
||||
masks = (hyper_in @ upscaled_embedding.view(b, c, h * w)).view(b, -1, h, w)
|
||||
# Generate mask quality predictions
|
||||
iou_pred = self.iou_prediction_head(iou_token_out)
|
||||
return masks, iou_pred
|
||||
@@ -0,0 +1,257 @@
|
||||
# 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 math
|
||||
from typing import List, Optional, Tuple, Type
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class LayerNorm2d(nn.Module):
|
||||
def __init__(self, num_channels: int, eps: float = 1e-6) -> None:
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.ones(num_channels))
|
||||
self.bias = nn.Parameter(torch.zeros(num_channels))
|
||||
self.eps = eps
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
u = x.mean(1, keepdim=True)
|
||||
s = (x - u).pow(2).mean(1, keepdim=True)
|
||||
x = (x - u) / torch.sqrt(s + self.eps)
|
||||
x = self.weight[:, None, None] * x + self.bias[:, None, None]
|
||||
return x
|
||||
|
||||
|
||||
class PatchEmbed(nn.Module):
|
||||
"""2D Image to Patch Embedding"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
img_size,
|
||||
patch_size,
|
||||
in_chans,
|
||||
embed_dim,
|
||||
):
|
||||
super().__init__()
|
||||
self.proj = nn.Conv2d(
|
||||
in_chans,
|
||||
embed_dim,
|
||||
kernel_size=(patch_size, patch_size),
|
||||
stride=(patch_size, patch_size),
|
||||
bias=True,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
B, C, H, W = x.shape
|
||||
x = self.proj(x)
|
||||
return x
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
num_heads,
|
||||
qkv_bias,
|
||||
qk_scale=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.num_heads = num_heads
|
||||
head_dim = dim // num_heads
|
||||
self.scale = qk_scale or head_dim**-0.5
|
||||
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
|
||||
self.proj = nn.Linear(dim, dim)
|
||||
|
||||
def forward(self, x):
|
||||
B, N, C = x.shape
|
||||
qkv = (
|
||||
self.qkv(x)
|
||||
.reshape(B, N, 3, self.num_heads, C // self.num_heads)
|
||||
.permute(2, 0, 3, 1, 4)
|
||||
)
|
||||
q, k, v = (
|
||||
qkv[0],
|
||||
qkv[1],
|
||||
qkv[2],
|
||||
)
|
||||
attn = (q @ k.transpose(-2, -1)) * self.scale
|
||||
attn = attn.softmax(dim=-1)
|
||||
x = (attn @ v).transpose(1, 2).reshape(B, N, C)
|
||||
x = self.proj(x)
|
||||
return x
|
||||
|
||||
|
||||
class Mlp(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_features,
|
||||
hidden_features=None,
|
||||
out_features=None,
|
||||
act_layer=nn.GELU,
|
||||
):
|
||||
super().__init__()
|
||||
out_features = out_features or in_features
|
||||
hidden_features = hidden_features or in_features
|
||||
self.fc1 = nn.Linear(in_features, hidden_features)
|
||||
self.act = act_layer()
|
||||
self.fc2 = nn.Linear(hidden_features, out_features)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.fc1(x)
|
||||
x = self.act(x)
|
||||
x = self.fc2(x)
|
||||
return x
|
||||
|
||||
|
||||
class Block(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
num_heads,
|
||||
mlp_ratio=4.0,
|
||||
qkv_bias=False,
|
||||
qk_scale=None,
|
||||
act_layer=nn.GELU,
|
||||
):
|
||||
super().__init__()
|
||||
self.norm1 = nn.LayerNorm(dim, eps=1e-6)
|
||||
self.attn = Attention(
|
||||
dim,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
)
|
||||
self.norm2 = nn.LayerNorm(dim, eps=1e-6)
|
||||
mlp_hidden_dim = int(dim * mlp_ratio)
|
||||
self.mlp = Mlp(
|
||||
in_features=dim,
|
||||
hidden_features=mlp_hidden_dim,
|
||||
act_layer=act_layer,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
x = x + self.attn(self.norm1(x))
|
||||
x = x + self.mlp(self.norm2(x))
|
||||
return x
|
||||
|
||||
|
||||
@torch.jit.export
|
||||
def get_abs_pos(
|
||||
abs_pos: torch.Tensor, has_cls_token: bool, hw: List[int]
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Calculate absolute positional embeddings. If needed, resize embeddings and remove cls_token
|
||||
dimension for the original embeddings.
|
||||
Args:
|
||||
abs_pos (Tensor): absolute positional embeddings with (1, num_position, C).
|
||||
has_cls_token (bool): If true, has 1 embedding in abs_pos for cls token.
|
||||
hw (Tuple): size of input image tokens.
|
||||
|
||||
Returns:
|
||||
Absolute positional embeddings after processing with shape (1, H, W, C)
|
||||
"""
|
||||
h = hw[0]
|
||||
w = hw[1]
|
||||
if has_cls_token:
|
||||
abs_pos = abs_pos[:, 1:]
|
||||
xy_num = abs_pos.shape[1]
|
||||
size = int(math.sqrt(xy_num))
|
||||
assert size * size == xy_num
|
||||
|
||||
if size != h or size != w:
|
||||
new_abs_pos = F.interpolate(
|
||||
abs_pos.reshape(1, size, size, -1).permute(0, 3, 1, 2),
|
||||
size=(h, w),
|
||||
mode="bicubic",
|
||||
align_corners=False,
|
||||
)
|
||||
return new_abs_pos.permute(0, 2, 3, 1)
|
||||
else:
|
||||
return abs_pos.reshape(1, h, w, -1)
|
||||
|
||||
|
||||
# Image encoder for efficient SAM.
|
||||
class ImageEncoderViT(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
img_size: int,
|
||||
patch_size: int,
|
||||
in_chans: int,
|
||||
patch_embed_dim: int,
|
||||
normalization_type: str,
|
||||
depth: int,
|
||||
num_heads: int,
|
||||
mlp_ratio: float,
|
||||
neck_dims: List[int],
|
||||
act_layer: Type[nn.Module],
|
||||
) -> None:
|
||||
"""
|
||||
Args:
|
||||
img_size (int): Input image size.
|
||||
patch_size (int): Patch size.
|
||||
in_chans (int): Number of input image channels.
|
||||
patch_embed_dim (int): Patch embedding dimension.
|
||||
depth (int): Depth of ViT.
|
||||
num_heads (int): Number of attention heads in each ViT block.
|
||||
mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.
|
||||
act_layer (nn.Module): Activation layer.
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
self.img_size = img_size
|
||||
self.image_embedding_size = img_size // ((patch_size if patch_size > 0 else 1))
|
||||
self.transformer_output_dim = ([patch_embed_dim] + neck_dims)[-1]
|
||||
self.pretrain_use_cls_token = True
|
||||
pretrain_img_size = 224
|
||||
self.patch_embed = PatchEmbed(img_size, patch_size, in_chans, patch_embed_dim)
|
||||
# Initialize absolute positional embedding with pretrain image size.
|
||||
num_patches = (pretrain_img_size // patch_size) * (
|
||||
pretrain_img_size // patch_size
|
||||
)
|
||||
num_positions = num_patches + 1
|
||||
self.pos_embed = nn.Parameter(torch.zeros(1, num_positions, patch_embed_dim))
|
||||
self.blocks = nn.ModuleList()
|
||||
for i in range(depth):
|
||||
vit_block = Block(patch_embed_dim, num_heads, mlp_ratio, True)
|
||||
self.blocks.append(vit_block)
|
||||
self.neck = nn.Sequential(
|
||||
nn.Conv2d(
|
||||
patch_embed_dim,
|
||||
neck_dims[0],
|
||||
kernel_size=1,
|
||||
bias=False,
|
||||
),
|
||||
LayerNorm2d(neck_dims[0]),
|
||||
nn.Conv2d(
|
||||
neck_dims[0],
|
||||
neck_dims[0],
|
||||
kernel_size=3,
|
||||
padding=1,
|
||||
bias=False,
|
||||
),
|
||||
LayerNorm2d(neck_dims[0]),
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
assert (
|
||||
x.shape[2] == self.img_size and x.shape[3] == self.img_size
|
||||
), "input image size must match self.img_size"
|
||||
x = self.patch_embed(x)
|
||||
# B C H W -> B H W C
|
||||
x = x.permute(0, 2, 3, 1)
|
||||
x = x + get_abs_pos(
|
||||
self.pos_embed, self.pretrain_use_cls_token, [x.shape[1], x.shape[2]]
|
||||
)
|
||||
num_patches = x.shape[1]
|
||||
assert x.shape[2] == num_patches
|
||||
x = x.reshape(x.shape[0], num_patches * num_patches, x.shape[3])
|
||||
for blk in self.blocks:
|
||||
x = blk(x)
|
||||
x = x.reshape(x.shape[0], num_patches, num_patches, x.shape[2])
|
||||
x = self.neck(x.permute(0, 3, 1, 2))
|
||||
return x
|
||||
@@ -0,0 +1,29 @@
|
||||
from typing import Type
|
||||
|
||||
from torch import nn
|
||||
|
||||
|
||||
# Lightly adapted from
|
||||
# https://github.com/facebookresearch/MaskFormer/blob/main/mask_former/modeling/transformer/transformer_predictor.py # noqa
|
||||
class MLPBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_dim: int,
|
||||
hidden_dim: int,
|
||||
output_dim: int,
|
||||
num_layers: int,
|
||||
act: Type[nn.Module],
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.num_layers = num_layers
|
||||
h = [hidden_dim] * (num_layers - 1)
|
||||
self.layers = nn.ModuleList(
|
||||
nn.Sequential(nn.Linear(n, k), act())
|
||||
for n, k in zip([input_dim] + h, [hidden_dim] * num_layers)
|
||||
)
|
||||
self.fc = nn.Linear(hidden_dim, output_dim)
|
||||
|
||||
def forward(self, x):
|
||||
for layer in self.layers:
|
||||
x = layer(x)
|
||||
return self.fc(x)
|
||||
@@ -0,0 +1,266 @@
|
||||
import math
|
||||
from typing import Tuple, Type
|
||||
import torch
|
||||
from torch import nn, Tensor
|
||||
from .mlp import MLPBlock
|
||||
|
||||
|
||||
|
||||
|
||||
class TwoWayTransformer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
depth: int,
|
||||
embedding_dim: int,
|
||||
num_heads: int,
|
||||
mlp_dim: int,
|
||||
activation: Type[nn.Module],
|
||||
normalize_before_activation: bool,
|
||||
attention_downsample_rate: int = 2,
|
||||
) -> None:
|
||||
"""
|
||||
A transformer decoder that attends to an input image using
|
||||
queries whose positional embedding is supplied.
|
||||
|
||||
Args:
|
||||
depth (int): number of layers in the transformer
|
||||
embedding_dim (int): the channel dimension for the input embeddings
|
||||
num_heads (int): the number of heads for multihead attention. Must
|
||||
divide embedding_dim
|
||||
mlp_dim (int): the channel dimension internal to the MLP block
|
||||
activation (nn.Module): the activation to use in the MLP block
|
||||
"""
|
||||
super().__init__()
|
||||
self.depth = depth
|
||||
self.embedding_dim = embedding_dim
|
||||
self.num_heads = num_heads
|
||||
self.mlp_dim = mlp_dim
|
||||
self.layers = nn.ModuleList()
|
||||
|
||||
for i in range(depth):
|
||||
curr_layer = TwoWayAttentionBlock(
|
||||
embedding_dim=embedding_dim,
|
||||
num_heads=num_heads,
|
||||
mlp_dim=mlp_dim,
|
||||
activation=activation,
|
||||
normalize_before_activation=normalize_before_activation,
|
||||
attention_downsample_rate=attention_downsample_rate,
|
||||
skip_first_layer_pe=(i == 0),
|
||||
)
|
||||
self.layers.append(curr_layer)
|
||||
|
||||
self.final_attn_token_to_image = AttentionForTwoWayAttentionBlock(
|
||||
embedding_dim,
|
||||
num_heads,
|
||||
downsample_rate=attention_downsample_rate,
|
||||
)
|
||||
self.norm_final_attn = nn.LayerNorm(embedding_dim)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
image_embedding: Tensor,
|
||||
image_pe: Tensor,
|
||||
point_embedding: Tensor,
|
||||
) -> Tuple[Tensor, Tensor]:
|
||||
"""
|
||||
Args:
|
||||
image_embedding (torch.Tensor): image to attend to. Should be shape
|
||||
B x embedding_dim x h x w for any h and w.
|
||||
image_pe (torch.Tensor): the positional encoding to add to the image. Must
|
||||
have the same shape as image_embedding.
|
||||
point_embedding (torch.Tensor): the embedding to add to the query points.
|
||||
Must have shape B x N_points x embedding_dim for any N_points.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: the processed point_embedding
|
||||
torch.Tensor: the processed image_embedding
|
||||
"""
|
||||
|
||||
# BxCxHxW -> BxHWxC == B x N_image_tokens x C
|
||||
bs, c, h, w = image_embedding.shape
|
||||
image_embedding = image_embedding.flatten(2).permute(0, 2, 1)
|
||||
image_pe = image_pe.flatten(2).permute(0, 2, 1)
|
||||
|
||||
# Prepare queries
|
||||
queries = point_embedding
|
||||
keys = image_embedding
|
||||
|
||||
# Apply transformer blocks and final layernorm
|
||||
for idx, layer in enumerate(self.layers):
|
||||
queries, keys = layer(
|
||||
queries=queries,
|
||||
keys=keys,
|
||||
query_pe=point_embedding,
|
||||
key_pe=image_pe,
|
||||
)
|
||||
|
||||
# Apply the final attention layer from the points to the image
|
||||
q = queries + point_embedding
|
||||
k = keys + image_pe
|
||||
attn_out = self.final_attn_token_to_image(q=q, k=k, v=keys)
|
||||
queries = queries + attn_out
|
||||
queries = self.norm_final_attn(queries)
|
||||
return queries, keys
|
||||
|
||||
|
||||
class TwoWayAttentionBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
embedding_dim: int,
|
||||
num_heads: int,
|
||||
mlp_dim: int,
|
||||
activation: Type[nn.Module],
|
||||
normalize_before_activation: bool,
|
||||
attention_downsample_rate: int = 2,
|
||||
skip_first_layer_pe: bool = False,
|
||||
) -> None:
|
||||
"""
|
||||
A transformer block with four layers: (1) self-attention of sparse
|
||||
inputs, (2) cross attention of sparse inputs to dense inputs, (3) mlp
|
||||
block on sparse inputs, and (4) cross attention of dense inputs to sparse
|
||||
inputs.
|
||||
|
||||
Arguments:
|
||||
embedding_dim (int): the channel dimension of the embeddings
|
||||
num_heads (int): the number of heads in the attention layers
|
||||
mlp_dim (int): the hidden dimension of the mlp block
|
||||
activation (nn.Module): the activation of the mlp block
|
||||
skip_first_layer_pe (bool): skip the PE on the first layer
|
||||
"""
|
||||
super().__init__()
|
||||
self.self_attn = AttentionForTwoWayAttentionBlock(embedding_dim, num_heads)
|
||||
self.norm1 = nn.LayerNorm(embedding_dim)
|
||||
|
||||
self.cross_attn_token_to_image = AttentionForTwoWayAttentionBlock(
|
||||
embedding_dim,
|
||||
num_heads,
|
||||
downsample_rate=attention_downsample_rate,
|
||||
)
|
||||
self.norm2 = nn.LayerNorm(embedding_dim)
|
||||
|
||||
self.mlp = MLPBlock(
|
||||
embedding_dim,
|
||||
mlp_dim,
|
||||
embedding_dim,
|
||||
1,
|
||||
activation,
|
||||
)
|
||||
|
||||
self.norm3 = nn.LayerNorm(embedding_dim)
|
||||
|
||||
self.norm4 = nn.LayerNorm(embedding_dim)
|
||||
self.cross_attn_image_to_token = AttentionForTwoWayAttentionBlock(
|
||||
embedding_dim,
|
||||
num_heads,
|
||||
downsample_rate=attention_downsample_rate,
|
||||
)
|
||||
|
||||
self.skip_first_layer_pe = skip_first_layer_pe
|
||||
|
||||
def forward(
|
||||
self, queries: Tensor, keys: Tensor, query_pe: Tensor, key_pe: Tensor
|
||||
) -> Tuple[Tensor, Tensor]:
|
||||
# Self attention block
|
||||
if not self.skip_first_layer_pe:
|
||||
queries = queries + query_pe
|
||||
attn_out = self.self_attn(q=queries, k=queries, v=queries)
|
||||
queries = queries + attn_out
|
||||
queries = self.norm1(queries)
|
||||
|
||||
# Cross attention block, tokens attending to image embedding
|
||||
q = queries + query_pe
|
||||
k = keys + key_pe
|
||||
attn_out = self.cross_attn_token_to_image(q=q, k=k, v=keys)
|
||||
queries = queries + attn_out
|
||||
queries = self.norm2(queries)
|
||||
|
||||
# MLP block
|
||||
mlp_out = self.mlp(queries)
|
||||
queries = queries + mlp_out
|
||||
queries = self.norm3(queries)
|
||||
|
||||
# Cross attention block, image embedding attending to tokens
|
||||
q = queries + query_pe
|
||||
k = keys + key_pe
|
||||
attn_out = self.cross_attn_image_to_token(q=k, k=q, v=queries)
|
||||
keys = keys + attn_out
|
||||
keys = self.norm4(keys)
|
||||
|
||||
return queries, keys
|
||||
|
||||
|
||||
class AttentionForTwoWayAttentionBlock(nn.Module):
|
||||
"""
|
||||
An attention layer that allows for downscaling the size of the embedding
|
||||
after projection to queries, keys, and values.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
embedding_dim: int,
|
||||
num_heads: int,
|
||||
downsample_rate: int = 1,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.embedding_dim = embedding_dim
|
||||
self.internal_dim = embedding_dim // downsample_rate
|
||||
self.num_heads = num_heads
|
||||
assert (
|
||||
self.internal_dim % num_heads == 0
|
||||
), "num_heads must divide embedding_dim."
|
||||
self.c_per_head = self.internal_dim / num_heads
|
||||
self.inv_sqrt_c_per_head = 1.0 / math.sqrt(self.c_per_head)
|
||||
|
||||
self.q_proj = nn.Linear(embedding_dim, self.internal_dim)
|
||||
self.k_proj = nn.Linear(embedding_dim, self.internal_dim)
|
||||
self.v_proj = nn.Linear(embedding_dim, self.internal_dim)
|
||||
self.out_proj = nn.Linear(self.internal_dim, embedding_dim)
|
||||
self._reset_parameters()
|
||||
|
||||
def _reset_parameters(self) -> None:
|
||||
# The fan_out is incorrect, but matches pytorch's initialization
|
||||
# for which qkv is a single 3*embedding_dim x embedding_dim matrix
|
||||
fan_in = self.embedding_dim
|
||||
fan_out = 3 * self.internal_dim
|
||||
# Xavier uniform with our custom fan_out
|
||||
bnd = math.sqrt(6 / (fan_in + fan_out))
|
||||
nn.init.uniform_(self.q_proj.weight, -bnd, bnd)
|
||||
nn.init.uniform_(self.k_proj.weight, -bnd, bnd)
|
||||
nn.init.uniform_(self.v_proj.weight, -bnd, bnd)
|
||||
# out_proj.weight is left with default initialization, like pytorch attention
|
||||
nn.init.zeros_(self.q_proj.bias)
|
||||
nn.init.zeros_(self.k_proj.bias)
|
||||
nn.init.zeros_(self.v_proj.bias)
|
||||
nn.init.zeros_(self.out_proj.bias)
|
||||
|
||||
def _separate_heads(self, x: Tensor, num_heads: int) -> Tensor:
|
||||
b, n, c = x.shape
|
||||
x = x.reshape(b, n, num_heads, c // num_heads)
|
||||
return x.transpose(1, 2) # B x N_heads x N_tokens x C_per_head
|
||||
|
||||
def _recombine_heads(self, x: Tensor) -> Tensor:
|
||||
b, n_heads, n_tokens, c_per_head = x.shape
|
||||
x = x.transpose(1, 2)
|
||||
return x.reshape(b, n_tokens, n_heads * c_per_head) # B x N_tokens x C
|
||||
|
||||
def forward(self, q: Tensor, k: Tensor, v: Tensor) -> Tensor:
|
||||
# Input projections
|
||||
q = self.q_proj(q)
|
||||
k = self.k_proj(k)
|
||||
v = self.v_proj(v)
|
||||
|
||||
# Separate into heads
|
||||
q = self._separate_heads(q, self.num_heads)
|
||||
k = self._separate_heads(k, self.num_heads)
|
||||
v = self._separate_heads(v, self.num_heads)
|
||||
|
||||
# Attention
|
||||
_, _, _, c_per_head = q.shape
|
||||
attn = q @ k.permute(0, 1, 3, 2) # B x N_heads x N_tokens x N_tokens
|
||||
attn = attn * self.inv_sqrt_c_per_head
|
||||
attn = torch.softmax(attn, dim=-1)
|
||||
# Get output
|
||||
out = attn @ v
|
||||
out = self._recombine_heads(out)
|
||||
out = self.out_proj(out)
|
||||
return out
|
||||
@@ -0,0 +1,113 @@
|
||||
# coding=utf-8
|
||||
# Copyright 2022 EleutherAI and the HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX
|
||||
# and OPT implementations in this library. It has been modified from its
|
||||
# original forms to accommodate minor architectural differences compared
|
||||
# to GPT-NeoX and OPT used by the Meta AI team that trained the model.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
""" Evf model configuration"""
|
||||
|
||||
from transformers.configuration_utils import PretrainedConfig
|
||||
from transformers.utils import logging
|
||||
|
||||
logger = logging.get_logger(__name__)
|
||||
|
||||
EVF_PRETRAINED_CONFIG_ARCHIVE_MAP = {}
|
||||
|
||||
|
||||
class EvfConfig(PretrainedConfig):
|
||||
r"""
|
||||
This is the configuration class to store the configuration of a [`EvfSam`].
|
||||
|
||||
Configuration objects inherit from [`PretrainedConfig`] and can be used to control the model outputs. Read the
|
||||
documentation from [`PretrainedConfig`] for more information.
|
||||
|
||||
Args:
|
||||
hidden_size (`int`, *optional*, defaults to 4096):
|
||||
Dimension of the hidden representations.
|
||||
pretraining_tp (`int`, *optional*, defaults to `1`):
|
||||
Experimental feature. Tensor parallelism rank used during pretraining. Please refer to [this
|
||||
document](https://huggingface.co/docs/transformers/parallelism) to understand more about it. This value is
|
||||
necessary to ensure exact reproducibility of the pretraining results. Please refer to [this
|
||||
issue](https://github.com/pytorch/pytorch/issues/76232).
|
||||
rope_scaling (`Dict`, *optional*):
|
||||
Dictionary containing the scaling configuration for the RoPE embeddings. Currently supports three scaling
|
||||
strategies: linear and dynamic. Their scaling factor must be an float greater than 1. The expected format
|
||||
is `{"type": strategy name, "factor": scaling factor}`. When using this flag, don't update
|
||||
`max_position_embeddings` to the expected new maximum. See the following thread for more information on how
|
||||
these scaling strategies behave:
|
||||
https://www.reddit.com/r/LocalLLaMA/comments/14mrgpr/dynamically_scaled_rope_further_increases/. This is an
|
||||
experimental feature, subject to breaking API changes in future versions.
|
||||
|
||||
Example:
|
||||
|
||||
```python
|
||||
|
||||
>>> configuration = EvfConfig()
|
||||
>>> model = EvfSam(configuration)
|
||||
|
||||
>>> # Accessing the model configuration
|
||||
>>> configuration = model.config
|
||||
```"""
|
||||
model_type = "evf"
|
||||
keys_to_ignore_at_inference = ["past_key_values"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size=768,
|
||||
pad_token_id=1,
|
||||
bos_token_id=0,
|
||||
eos_token_id=2,
|
||||
pretraining_tp=1,
|
||||
tie_word_embeddings=False,
|
||||
rope_scaling=None,
|
||||
out_dim=256,
|
||||
**kwargs,
|
||||
):
|
||||
self.hidden_size = hidden_size
|
||||
self.out_dim = out_dim
|
||||
|
||||
# self.pretraining_tp = pretraining_tp
|
||||
# self.rope_scaling = rope_scaling
|
||||
# self._rope_scaling_validation()
|
||||
|
||||
super().__init__(
|
||||
pad_token_id=pad_token_id,
|
||||
bos_token_id=bos_token_id,
|
||||
eos_token_id=eos_token_id,
|
||||
tie_word_embeddings=tie_word_embeddings,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def _rope_scaling_validation(self):
|
||||
"""
|
||||
Validate the `rope_scaling` configuration.
|
||||
"""
|
||||
if self.rope_scaling is None:
|
||||
return
|
||||
|
||||
if not isinstance(self.rope_scaling, dict) or len(self.rope_scaling) != 2:
|
||||
raise ValueError(
|
||||
"`rope_scaling` must be a dictionary with with two fields, `name` and `factor`, "
|
||||
f"got {self.rope_scaling}"
|
||||
)
|
||||
rope_scaling_type = self.rope_scaling.get("type", None)
|
||||
rope_scaling_factor = self.rope_scaling.get("factor", None)
|
||||
if rope_scaling_type is None or rope_scaling_type not in ["linear", "dynamic"]:
|
||||
raise ValueError(
|
||||
f"`rope_scaling`'s name field must be one of ['linear', 'dynamic'], got {rope_scaling_type}"
|
||||
)
|
||||
if rope_scaling_factor is None or not isinstance(rope_scaling_factor, float) or rope_scaling_factor <= 1.0:
|
||||
raise ValueError(f"`rope_scaling`'s factor field must be an float > 1, got {rope_scaling_factor}")
|
||||
@@ -0,0 +1,313 @@
|
||||
from typing import List, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from transformers import PreTrainedModel, AutoConfig, AutoModelForCausalLM
|
||||
from .EfficientSAM.efficient_sam.build_efficient_sam import build_efficient_sam_vits, build_efficient_sam_vitt
|
||||
from .unilm.beit3.modeling_utils import BEiT3Wrapper, _get_base_config, _get_large_config
|
||||
from .configuration_evf import EvfConfig
|
||||
|
||||
|
||||
def dice_loss(
|
||||
inputs: torch.Tensor,
|
||||
targets: torch.Tensor,
|
||||
num_masks: float,
|
||||
scale=1000, # 100000.0,
|
||||
eps=1e-6,
|
||||
):
|
||||
"""
|
||||
Compute the DICE loss, similar to generalized IOU for masks
|
||||
Args:
|
||||
inputs: A float tensor of arbitrary shape.
|
||||
The predictions for each example.
|
||||
targets: A float tensor with the same shape as inputs. Stores the binary
|
||||
classification label for each element in inputs
|
||||
(0 for the negative class and 1 for the positive class).
|
||||
"""
|
||||
inputs = inputs.sigmoid()
|
||||
inputs = inputs.flatten(1, 2)
|
||||
targets = targets.flatten(1, 2)
|
||||
numerator = 2 * (inputs / scale * targets).sum(-1)
|
||||
denominator = (inputs / scale).sum(-1) + (targets / scale).sum(-1)
|
||||
loss = 1 - (numerator + eps) / (denominator + eps)
|
||||
loss = loss.sum() / (num_masks + 1e-8)
|
||||
return loss
|
||||
|
||||
|
||||
def sigmoid_ce_loss(
|
||||
inputs: torch.Tensor,
|
||||
targets: torch.Tensor,
|
||||
num_masks: float,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
inputs: A float tensor of arbitrary shape.
|
||||
The predictions for each example.
|
||||
targets: A float tensor with the same shape as inputs. Stores the binary
|
||||
classification label for each element in inputs
|
||||
(0 for the negative class and 1 for the positive class).
|
||||
Returns:
|
||||
Loss tensor
|
||||
"""
|
||||
loss = F.binary_cross_entropy_with_logits(inputs, targets, reduction="none")
|
||||
loss = loss.flatten(1, 2).mean(1).sum() / (num_masks + 1e-8)
|
||||
return loss
|
||||
|
||||
|
||||
|
||||
class EvfEffiSamModel(PreTrainedModel):
|
||||
config_class = EvfConfig
|
||||
def __init__(
|
||||
self,
|
||||
config,
|
||||
**kwargs
|
||||
):
|
||||
super(EvfEffiSamModel, self).__init__(config)
|
||||
|
||||
self.config = config
|
||||
self.vision_pretrained = kwargs.get("vision_pretrained", None)
|
||||
self.encoder_pretrained = kwargs.get("encoder_pretrained", None)
|
||||
self.dice_loss_weight = kwargs.get("dice_loss_weight", None)
|
||||
self.bce_loss_weight = kwargs.get("bce_loss_weight", None)
|
||||
self.train_mask_decoder = kwargs.get("train_mask_decoder", False)
|
||||
self.initialize_evf_modules(config)
|
||||
|
||||
|
||||
def initialize_evf_modules(self, config):
|
||||
# EffiSAM
|
||||
if config.sam_scale=="tiny":
|
||||
self.visual_model = build_efficient_sam_vitt(self.vision_pretrained)
|
||||
elif config.sam_scale=="small":
|
||||
# vits scale, or without pretrained weight (self.vision_pretrained=None)
|
||||
self.visual_model = build_efficient_sam_vits(self.vision_pretrained)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
for param in self.visual_model.parameters():
|
||||
param.requires_grad = False
|
||||
if self.train_mask_decoder:
|
||||
self.visual_model.mask_decoder.train()
|
||||
for param in self.visual_model.mask_decoder.parameters():
|
||||
param.requires_grad = True
|
||||
|
||||
# beit-3
|
||||
if self.config.mm_extractor_scale == "base":
|
||||
beit_config = _get_base_config()
|
||||
elif self.config.mm_extractor_scale == "large":
|
||||
beit_config = _get_large_config()
|
||||
else:
|
||||
raise AttributeError(f"model config should contain key 'mm_extractor_scale', with value 'base' or 'large'.")
|
||||
|
||||
self.mm_extractor = BEiT3Wrapper(beit_config)
|
||||
if self.encoder_pretrained is not None:
|
||||
beit_state_dict = torch.load(self.encoder_pretrained)["model"]
|
||||
self.mm_extractor.load_state_dict(
|
||||
beit_state_dict,
|
||||
strict=False
|
||||
)
|
||||
|
||||
for param in self.mm_extractor.parameters():
|
||||
param.requires_grad = True
|
||||
|
||||
# Projection layer
|
||||
in_dim = config.hidden_size
|
||||
assert in_dim==beit_config.encoder_embed_dim, \
|
||||
f"projection layer dim {in_dim} mismatch with mm_extractor dim {beit_config.encoder_embed_dim}"
|
||||
out_dim = config.out_dim
|
||||
text_fc = [
|
||||
nn.Linear(in_dim, in_dim),
|
||||
nn.ReLU(),
|
||||
nn.Linear(in_dim, out_dim)
|
||||
]
|
||||
self.text_hidden_fcs = nn.ModuleList([nn.Sequential(*text_fc)])
|
||||
self.text_hidden_fcs.train()
|
||||
for param in self.text_hidden_fcs.parameters():
|
||||
param.requires_grad = True
|
||||
|
||||
def get_visual_embs(self, pixel_values: torch.Tensor):
|
||||
with torch.no_grad():
|
||||
image_embeddings_list = []
|
||||
for i in range(pixel_values.shape[0]):
|
||||
torch.cuda.empty_cache()
|
||||
image_embeddings = self.visual_model.image_encoder(
|
||||
pixel_values[i].unsqueeze(0)
|
||||
)
|
||||
image_embeddings_list.append(image_embeddings)
|
||||
torch.cuda.empty_cache()
|
||||
image_embeddings = torch.cat(image_embeddings_list, 0)
|
||||
return image_embeddings
|
||||
|
||||
def forward(
|
||||
self,
|
||||
images: torch.Tensor,
|
||||
images_evf: torch.Tensor,
|
||||
input_ids: torch.Tensor,
|
||||
attention_masks: torch.Tensor,
|
||||
offset: torch.Tensor,
|
||||
masks_list: List[torch.Tensor],
|
||||
label_list: List[torch.Tensor],
|
||||
resize_list: List[tuple],
|
||||
inference: bool = False,
|
||||
**kwargs,
|
||||
):
|
||||
image_embeddings = self.get_visual_embs(images)
|
||||
batch_size = image_embeddings.shape[0]
|
||||
assert batch_size == len(offset) - 1
|
||||
|
||||
images_evf_list = []
|
||||
for i in range(len(offset) - 1):
|
||||
start_i, end_i = offset[i], offset[i + 1]
|
||||
images_evf_i = (
|
||||
images_evf[i]
|
||||
.unsqueeze(0)
|
||||
.expand(end_i - start_i, -1, -1, -1)
|
||||
.contiguous()
|
||||
)
|
||||
images_evf_list.append(images_evf_i)
|
||||
images_evf = torch.cat(images_evf_list, dim=0)
|
||||
|
||||
multimask_output = False
|
||||
output = self.mm_extractor.beit3(
|
||||
visual_tokens=images_evf,
|
||||
textual_tokens=input_ids,
|
||||
text_padding_position=~attention_masks
|
||||
)
|
||||
|
||||
feat = output["encoder_out"][:, :1, ...]
|
||||
|
||||
feat = self.text_hidden_fcs[0](feat)
|
||||
feat = torch.split(feat, [offset[i+1] - offset[i] for i in range(len(offset)-1)])
|
||||
|
||||
pred_masks = []
|
||||
for i in range(len(feat)):
|
||||
sparse_embeddings = feat[i].unsqueeze(0)
|
||||
sparse_embeddings = sparse_embeddings.to(feat[i].dtype)
|
||||
low_res_masks, iou_predictions = self.visual_model.mask_decoder(
|
||||
image_embeddings=image_embeddings[i].unsqueeze(0),
|
||||
image_pe=self.visual_model.prompt_encoder.get_dense_pe(),
|
||||
sparse_prompt_embeddings=sparse_embeddings,
|
||||
multimask_output=multimask_output,
|
||||
)
|
||||
|
||||
if multimask_output:
|
||||
sorted_ids = torch.argsort(iou_predictions, dim=-1, descending=True)
|
||||
low_res_masks = torch.take_along_dim(low_res_masks, sorted_ids[..., None, None], dim=1)
|
||||
|
||||
pred_mask = self.postprocess_masks(
|
||||
low_res_masks[:, :1],
|
||||
input_size=resize_list[i],
|
||||
original_size=label_list[i].shape,
|
||||
)
|
||||
pred_masks.append(pred_mask[:, 0])
|
||||
|
||||
gt_masks = masks_list
|
||||
|
||||
if inference:
|
||||
return {
|
||||
"pred_masks": pred_masks,
|
||||
"gt_masks": gt_masks,
|
||||
}
|
||||
|
||||
mask_bce_loss = 0
|
||||
mask_dice_loss = 0
|
||||
num_masks = 0
|
||||
for batch_idx in range(len(pred_masks)):
|
||||
gt_mask = gt_masks[batch_idx]
|
||||
pred_mask = pred_masks[batch_idx]
|
||||
|
||||
assert (
|
||||
gt_mask.shape[0] == pred_mask.shape[0]
|
||||
), "gt_mask.shape: {}, pred_mask.shape: {}".format(
|
||||
gt_mask.shape, pred_mask.shape
|
||||
)
|
||||
mask_bce_loss += (
|
||||
sigmoid_ce_loss(pred_mask, gt_mask, num_masks=gt_mask.shape[0])
|
||||
* gt_mask.shape[0]
|
||||
)
|
||||
mask_dice_loss += (
|
||||
dice_loss(pred_mask, gt_mask, num_masks=gt_mask.shape[0])
|
||||
* gt_mask.shape[0]
|
||||
)
|
||||
num_masks += gt_mask.shape[0]
|
||||
|
||||
mask_bce_loss = self.bce_loss_weight * mask_bce_loss / (num_masks + 1e-8)
|
||||
mask_dice_loss = self.dice_loss_weight * mask_dice_loss / (num_masks + 1e-8)
|
||||
mask_loss = mask_bce_loss + mask_dice_loss
|
||||
|
||||
loss = mask_loss
|
||||
|
||||
return {
|
||||
"loss": loss,
|
||||
"mask_bce_loss": mask_bce_loss,
|
||||
"mask_dice_loss": mask_dice_loss,
|
||||
"mask_loss": mask_loss,
|
||||
}
|
||||
|
||||
def postprocess_masks(
|
||||
self,
|
||||
masks: torch.Tensor,
|
||||
input_size: Tuple[int, ...],
|
||||
original_size: Tuple[int, ...],
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
pre-process of Effi-SAM is different from SAM, where there is no padding,
|
||||
so cropping is not needed in post-process.
|
||||
"""
|
||||
|
||||
dtype = masks.dtype
|
||||
|
||||
# masks = F.interpolate(
|
||||
# masks.float(),
|
||||
# (1024, 1024),
|
||||
# mode="bilinear",
|
||||
# align_corners=False,
|
||||
# )
|
||||
# masks = masks.to(dtype)
|
||||
# masks = masks[..., : input_size[0], : input_size[1]]
|
||||
|
||||
masks = F.interpolate(
|
||||
masks, original_size, mode="bilinear", align_corners=False
|
||||
)
|
||||
masks = masks.to(dtype)
|
||||
return masks
|
||||
|
||||
def inference(
|
||||
self,
|
||||
images,
|
||||
images_evf,
|
||||
input_ids,
|
||||
resize_list,
|
||||
original_size_list,
|
||||
multimask_output=False,
|
||||
):
|
||||
with torch.no_grad():
|
||||
image_embeddings = self.visual_model.image_encoder(images)
|
||||
|
||||
output = self.mm_extractor.beit3(visual_tokens=images_evf, textual_tokens=input_ids, text_padding_position=torch.zeros_like(input_ids))
|
||||
|
||||
feat = output["encoder_out"][:, :1, ...]
|
||||
feat = self.text_hidden_fcs[0](feat)
|
||||
sparse_embeddings = feat.unsqueeze(0)
|
||||
sparse_embeddings = sparse_embeddings.to(feat.dtype)
|
||||
low_res_masks, iou_predictions = self.visual_model.mask_decoder(
|
||||
image_embeddings=image_embeddings,
|
||||
image_pe=self.visual_model.prompt_encoder.get_dense_pe(),
|
||||
sparse_prompt_embeddings=sparse_embeddings,
|
||||
multimask_output=multimask_output,
|
||||
)
|
||||
if multimask_output:
|
||||
sorted_ids = torch.argsort(iou_predictions, dim=-1, descending=True)
|
||||
low_res_masks = torch.take_along_dim(low_res_masks, sorted_ids[..., None, None], dim=1)
|
||||
|
||||
pred_mask = self.postprocess_masks(
|
||||
low_res_masks[:, :1],
|
||||
input_size=resize_list[0],
|
||||
original_size=original_size_list[0],
|
||||
)
|
||||
|
||||
return pred_mask[:, 0]
|
||||
|
||||
|
||||
AutoConfig.register("evf", EvfConfig)
|
||||
AutoModelForCausalLM.register(EvfConfig, EvfEffiSamModel)
|
||||
@@ -0,0 +1,303 @@
|
||||
from typing import List
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from transformers import PreTrainedModel, AutoConfig, AutoModelForCausalLM
|
||||
from .segment_anything import build_sam_vit_h
|
||||
from .unilm.beit3.modeling_utils import BEiT3Wrapper, _get_base_config, _get_large_config
|
||||
from .configuration_evf import EvfConfig
|
||||
|
||||
def dice_loss(
|
||||
inputs: torch.Tensor,
|
||||
targets: torch.Tensor,
|
||||
num_masks: float,
|
||||
scale=1000, # 100000.0,
|
||||
eps=1e-6,
|
||||
):
|
||||
"""
|
||||
Compute the DICE loss, similar to generalized IOU for masks
|
||||
Args:
|
||||
inputs: A float tensor of arbitrary shape.
|
||||
The predictions for each example.
|
||||
targets: A float tensor with the same shape as inputs. Stores the binary
|
||||
classification label for each element in inputs
|
||||
(0 for the negative class and 1 for the positive class).
|
||||
"""
|
||||
inputs = inputs.sigmoid()
|
||||
inputs = inputs.flatten(1, 2)
|
||||
targets = targets.flatten(1, 2)
|
||||
numerator = 2 * (inputs / scale * targets).sum(-1)
|
||||
denominator = (inputs / scale).sum(-1) + (targets / scale).sum(-1)
|
||||
loss = 1 - (numerator + eps) / (denominator + eps)
|
||||
loss = loss.sum() / (num_masks + 1e-8)
|
||||
return loss
|
||||
|
||||
|
||||
def sigmoid_ce_loss(
|
||||
inputs: torch.Tensor,
|
||||
targets: torch.Tensor,
|
||||
num_masks: float,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
inputs: A float tensor of arbitrary shape.
|
||||
The predictions for each example.
|
||||
targets: A float tensor with the same shape as inputs. Stores the binary
|
||||
classification label for each element in inputs
|
||||
(0 for the negative class and 1 for the positive class).
|
||||
Returns:
|
||||
Loss tensor
|
||||
"""
|
||||
loss = F.binary_cross_entropy_with_logits(inputs, targets, reduction="none")
|
||||
loss = loss.flatten(1, 2).mean(1).sum() / (num_masks + 1e-8)
|
||||
return loss
|
||||
|
||||
|
||||
|
||||
class EvfSamModel(PreTrainedModel):
|
||||
config_class = EvfConfig
|
||||
def __init__(
|
||||
self,
|
||||
config,
|
||||
**kwargs
|
||||
):
|
||||
super(EvfSamModel, self).__init__(config)
|
||||
|
||||
self.config = config
|
||||
self.vision_pretrained = kwargs.get("vision_pretrained", None)
|
||||
self.encoder_pretrained = kwargs.get("encoder_pretrained", None)
|
||||
self.dice_loss_weight = kwargs.get("dice_loss_weight", None)
|
||||
self.bce_loss_weight = kwargs.get("bce_loss_weight", None)
|
||||
self.train_mask_decoder = kwargs.get("train_mask_decoder", False)
|
||||
self.train_prompt_encoder = kwargs.get("train_prompt_encoder", False)
|
||||
self.initialize_evf_modules(config)
|
||||
|
||||
|
||||
def initialize_evf_modules(self, config):
|
||||
# SAM
|
||||
if config.sam_scale=="huge":
|
||||
self.visual_model = build_sam_vit_h(self.vision_pretrained)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
for param in self.visual_model.parameters():
|
||||
param.requires_grad = False
|
||||
if self.train_mask_decoder:
|
||||
self.visual_model.mask_decoder.train()
|
||||
for param in self.visual_model.mask_decoder.parameters():
|
||||
param.requires_grad = True
|
||||
if self.train_prompt_encoder:
|
||||
self.visual_model.prompt_encoder.no_mask_embed.requires_grad_(True)
|
||||
|
||||
# beit-3
|
||||
if self.config.mm_extractor_scale == "base":
|
||||
beit_config = _get_base_config()
|
||||
elif self.config.mm_extractor_scale == "large":
|
||||
beit_config = _get_large_config()
|
||||
else:
|
||||
raise AttributeError(f"model config should contain key 'mm_extractor_scale', with value 'base' or 'large'.")
|
||||
|
||||
self.mm_extractor = BEiT3Wrapper(beit_config)
|
||||
if self.encoder_pretrained is not None:
|
||||
beit_state_dict = torch.load(self.encoder_pretrained)["model"]
|
||||
self.mm_extractor.load_state_dict(
|
||||
beit_state_dict,
|
||||
strict=False
|
||||
)
|
||||
|
||||
for param in self.mm_extractor.parameters():
|
||||
param.requires_grad = True
|
||||
|
||||
# Projection layer
|
||||
in_dim = config.hidden_size
|
||||
assert in_dim==beit_config.encoder_embed_dim, \
|
||||
f"projection layer dim {in_dim} mismatch with mm_extractor dim {beit_config.encoder_embed_dim}"
|
||||
out_dim = config.out_dim
|
||||
text_fc = [
|
||||
nn.Linear(in_dim, in_dim),
|
||||
nn.ReLU(),
|
||||
nn.Linear(in_dim, out_dim)
|
||||
]
|
||||
self.text_hidden_fcs = nn.ModuleList([nn.Sequential(*text_fc)])
|
||||
self.text_hidden_fcs.train()
|
||||
for param in self.text_hidden_fcs.parameters():
|
||||
param.requires_grad = True
|
||||
|
||||
def get_visual_embs(self, pixel_values: torch.FloatTensor):
|
||||
with torch.no_grad():
|
||||
image_embeddings_list = []
|
||||
for i in range(pixel_values.shape[0]):
|
||||
torch.cuda.empty_cache()
|
||||
image_embeddings = self.visual_model.image_encoder(
|
||||
pixel_values[i].unsqueeze(0)
|
||||
)
|
||||
image_embeddings_list.append(image_embeddings)
|
||||
torch.cuda.empty_cache()
|
||||
image_embeddings = torch.cat(image_embeddings_list, 0)
|
||||
return image_embeddings
|
||||
|
||||
def forward(
|
||||
self,
|
||||
images: torch.FloatTensor,
|
||||
images_evf: torch.FloatTensor,
|
||||
input_ids: torch.LongTensor,
|
||||
attention_masks: torch.LongTensor,
|
||||
offset: torch.LongTensor,
|
||||
masks_list: List[torch.FloatTensor],
|
||||
label_list: List[torch.Tensor],
|
||||
resize_list: List[tuple],
|
||||
inference: bool = False,
|
||||
**kwargs,
|
||||
):
|
||||
image_embeddings = self.get_visual_embs(images)
|
||||
batch_size = image_embeddings.shape[0]
|
||||
assert batch_size == len(offset) - 1
|
||||
|
||||
images_evf_list = []
|
||||
for i in range(len(offset) - 1):
|
||||
start_i, end_i = offset[i], offset[i + 1]
|
||||
images_evf_i = (
|
||||
images_evf[i]
|
||||
.unsqueeze(0)
|
||||
.expand(end_i - start_i, -1, -1, -1)
|
||||
.contiguous()
|
||||
)
|
||||
images_evf_list.append(images_evf_i)
|
||||
images_evf = torch.cat(images_evf_list, dim=0)
|
||||
|
||||
multimask_output = False
|
||||
output = self.mm_extractor.beit3(
|
||||
visual_tokens=images_evf,
|
||||
textual_tokens=input_ids,
|
||||
text_padding_position=~attention_masks
|
||||
)
|
||||
|
||||
feat = output["encoder_out"][:, :1, ...]
|
||||
|
||||
feat = self.text_hidden_fcs[0](feat)
|
||||
feat = torch.split(feat, [offset[i+1] - offset[i] for i in range(len(offset)-1)])
|
||||
|
||||
pred_masks = []
|
||||
for i in range(len(feat)):
|
||||
(
|
||||
sparse_embeddings,
|
||||
dense_embeddings,
|
||||
) = self.visual_model.prompt_encoder(
|
||||
points=None,
|
||||
boxes=None,
|
||||
masks=None,
|
||||
text_embeds=feat[i],
|
||||
)
|
||||
sparse_embeddings = sparse_embeddings.to(feat[i].dtype)
|
||||
low_res_masks, iou_predictions = self.visual_model.mask_decoder(
|
||||
image_embeddings=image_embeddings[i].unsqueeze(0),
|
||||
image_pe=self.visual_model.prompt_encoder.get_dense_pe(),
|
||||
sparse_prompt_embeddings=sparse_embeddings,
|
||||
dense_prompt_embeddings=dense_embeddings,
|
||||
multimask_output=multimask_output,
|
||||
)
|
||||
|
||||
if multimask_output:
|
||||
sorted_ids = torch.argsort(iou_predictions, dim=-1, descending=True)
|
||||
low_res_masks = torch.take_along_dim(low_res_masks, sorted_ids[..., None, None], dim=1)[:, :1]
|
||||
|
||||
pred_mask = self.visual_model.postprocess_masks(
|
||||
low_res_masks,
|
||||
input_size=resize_list[i],
|
||||
original_size=label_list[i].shape,
|
||||
)
|
||||
pred_masks.append(pred_mask[:, 0])
|
||||
|
||||
gt_masks = masks_list
|
||||
|
||||
if inference:
|
||||
return {
|
||||
"pred_masks": pred_masks,
|
||||
"gt_masks": gt_masks,
|
||||
}
|
||||
|
||||
mask_bce_loss = 0
|
||||
mask_dice_loss = 0
|
||||
num_masks = 0
|
||||
for batch_idx in range(len(pred_masks)):
|
||||
gt_mask = gt_masks[batch_idx]
|
||||
pred_mask = pred_masks[batch_idx]
|
||||
|
||||
assert (
|
||||
gt_mask.shape[0] == pred_mask.shape[0]
|
||||
), "gt_mask.shape: {}, pred_mask.shape: {}".format(
|
||||
gt_mask.shape, pred_mask.shape
|
||||
)
|
||||
mask_bce_loss += (
|
||||
sigmoid_ce_loss(pred_mask, gt_mask, num_masks=gt_mask.shape[0])
|
||||
* gt_mask.shape[0]
|
||||
)
|
||||
mask_dice_loss += (
|
||||
dice_loss(pred_mask, gt_mask, num_masks=gt_mask.shape[0])
|
||||
* gt_mask.shape[0]
|
||||
)
|
||||
num_masks += gt_mask.shape[0]
|
||||
|
||||
mask_bce_loss = self.bce_loss_weight * mask_bce_loss / (num_masks + 1e-8)
|
||||
mask_dice_loss = self.dice_loss_weight * mask_dice_loss / (num_masks + 1e-8)
|
||||
mask_loss = mask_bce_loss + mask_dice_loss
|
||||
|
||||
loss = mask_loss
|
||||
|
||||
return {
|
||||
"loss": loss,
|
||||
"mask_bce_loss": mask_bce_loss,
|
||||
"mask_dice_loss": mask_dice_loss,
|
||||
"mask_loss": mask_loss,
|
||||
}
|
||||
|
||||
def inference(
|
||||
self,
|
||||
images,
|
||||
images_evf,
|
||||
input_ids,
|
||||
resize_list,
|
||||
original_size_list,
|
||||
multimask_output=False,
|
||||
):
|
||||
with torch.no_grad():
|
||||
image_embeddings = self.visual_model.image_encoder(images)
|
||||
multimask_output = multimask_output
|
||||
|
||||
output = self.mm_extractor.beit3(visual_tokens=images_evf, textual_tokens=input_ids, text_padding_position=torch.zeros_like(input_ids))
|
||||
|
||||
feat = output["encoder_out"][:, :1, ...]
|
||||
feat = self.text_hidden_fcs[0](feat)
|
||||
(
|
||||
sparse_embeddings,
|
||||
dense_embeddings,
|
||||
) = self.visual_model.prompt_encoder(
|
||||
points=None,
|
||||
boxes=None,
|
||||
masks=None,
|
||||
text_embeds=feat,
|
||||
)
|
||||
sparse_embeddings = sparse_embeddings.to(feat.dtype)
|
||||
low_res_masks, iou_predictions = self.visual_model.mask_decoder(
|
||||
image_embeddings=image_embeddings,
|
||||
image_pe=self.visual_model.prompt_encoder.get_dense_pe(),
|
||||
sparse_prompt_embeddings=sparse_embeddings,
|
||||
dense_prompt_embeddings=dense_embeddings,
|
||||
multimask_output=multimask_output,
|
||||
)
|
||||
if multimask_output:
|
||||
sorted_ids = torch.argsort(iou_predictions, dim=-1, descending=True)
|
||||
low_res_masks = torch.take_along_dim(low_res_masks, sorted_ids[..., None, None], dim=1)[:, :1]
|
||||
|
||||
pred_mask = self.visual_model.postprocess_masks(
|
||||
low_res_masks,
|
||||
input_size=resize_list[0],
|
||||
original_size=original_size_list[0],
|
||||
)
|
||||
|
||||
return pred_mask[:, 0]
|
||||
|
||||
|
||||
AutoConfig.register("evf", EvfConfig)
|
||||
AutoModelForCausalLM.register(EvfConfig, EvfSamModel)
|
||||
@@ -0,0 +1,341 @@
|
||||
from typing import List
|
||||
import os
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from transformers import PreTrainedModel, AutoConfig, AutoModelForCausalLM
|
||||
from .segment_anything_2.sam2.build_sam import build_sam2
|
||||
from .unilm.beit3.modeling_utils import BEiT3Wrapper, _get_base_config, _get_large_config
|
||||
from .configuration_evf import EvfConfig
|
||||
from .segment_anything_2.sam2.utils.misc import load_video_frames
|
||||
from collections import OrderedDict
|
||||
|
||||
|
||||
def dice_loss(
|
||||
inputs: torch.Tensor,
|
||||
targets: torch.Tensor,
|
||||
num_masks: float,
|
||||
scale=1000, # 100000.0,
|
||||
eps=1e-6,
|
||||
):
|
||||
"""
|
||||
Compute the DICE loss, similar to generalized IOU for masks
|
||||
Args:
|
||||
inputs: A float tensor of arbitrary shape.
|
||||
The predictions for each example.
|
||||
targets: A float tensor with the same shape as inputs. Stores the binary
|
||||
classification label for each element in inputs
|
||||
(0 for the negative class and 1 for the positive class).
|
||||
"""
|
||||
inputs = inputs.sigmoid()
|
||||
inputs = inputs.flatten(1, 2)
|
||||
targets = targets.flatten(1, 2)
|
||||
numerator = 2 * (inputs / scale * targets).sum(-1)
|
||||
denominator = (inputs / scale).sum(-1) + (targets / scale).sum(-1)
|
||||
loss = 1 - (numerator + eps) / (denominator + eps)
|
||||
loss = loss.sum() / (num_masks + 1e-8)
|
||||
return loss
|
||||
|
||||
|
||||
def sigmoid_ce_loss(
|
||||
inputs: torch.Tensor,
|
||||
targets: torch.Tensor,
|
||||
num_masks: float,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
inputs: A float tensor of arbitrary shape.
|
||||
The predictions for each example.
|
||||
targets: A float tensor with the same shape as inputs. Stores the binary
|
||||
classification label for each element in inputs
|
||||
(0 for the negative class and 1 for the positive class).
|
||||
Returns:
|
||||
Loss tensor
|
||||
"""
|
||||
loss = F.binary_cross_entropy_with_logits(inputs, targets, reduction="none")
|
||||
loss = loss.flatten(1, 2).mean(1).sum() / (num_masks + 1e-8)
|
||||
return loss
|
||||
|
||||
class EvfSam2Model(PreTrainedModel):
|
||||
config_class = EvfConfig
|
||||
def __init__(
|
||||
self,
|
||||
config,
|
||||
**kwargs
|
||||
):
|
||||
super(EvfSam2Model, self).__init__(config)
|
||||
|
||||
self.config = config
|
||||
self.vision_pretrained = kwargs.get("vision_pretrained", None)
|
||||
self.encoder_pretrained = kwargs.get("encoder_pretrained", None)
|
||||
self.dice_loss_weight = kwargs.get("dice_loss_weight", None)
|
||||
self.bce_loss_weight = kwargs.get("bce_loss_weight", None)
|
||||
self.train_mask_decoder = kwargs.get("train_mask_decoder", False)
|
||||
self.train_prompt_encoder = kwargs.get("train_prompt_encoder", False)
|
||||
self.initialize_evf_modules(config)
|
||||
self._bb_feat_sizes = [
|
||||
(256, 256),
|
||||
(128, 128),
|
||||
(64, 64),
|
||||
]
|
||||
|
||||
def initialize_evf_modules(self, config):
|
||||
# SAM
|
||||
if config.sam_scale=="large":
|
||||
self.visual_model = build_sam2("sam2_hiera_l.yaml", self.vision_pretrained, device=None)
|
||||
elif config.sam_scale=="tiny":
|
||||
self.visual_model = build_sam2("sam2_hiera_t.yaml", self.vision_pretrained, device=None)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
for param in self.visual_model.parameters():
|
||||
param.requires_grad = False
|
||||
if self.train_mask_decoder:
|
||||
self.visual_model.sam_mask_decoder.train()
|
||||
for param in self.visual_model.sam_mask_decoder.parameters():
|
||||
param.requires_grad = True
|
||||
if self.train_prompt_encoder:
|
||||
self.visual_model.sam_prompt_encoder.no_mask_embed.requires_grad_(True)
|
||||
|
||||
# beit-3
|
||||
if self.config.mm_extractor_scale == "base":
|
||||
beit_config = _get_base_config()
|
||||
elif self.config.mm_extractor_scale == "large":
|
||||
beit_config = _get_large_config()
|
||||
else:
|
||||
raise AttributeError(f"model config should contain key 'mm_extractor_scale', with value 'base' or 'large'.")
|
||||
|
||||
self.mm_extractor = BEiT3Wrapper(beit_config)
|
||||
if self.encoder_pretrained is not None:
|
||||
beit_state_dict = torch.load(self.encoder_pretrained)["model"]
|
||||
self.mm_extractor.load_state_dict(
|
||||
beit_state_dict,
|
||||
strict=False
|
||||
)
|
||||
|
||||
for param in self.mm_extractor.parameters():
|
||||
param.requires_grad = True
|
||||
|
||||
# Projection layer
|
||||
in_dim = config.hidden_size
|
||||
assert in_dim==beit_config.encoder_embed_dim, \
|
||||
f"projection layer dim {in_dim} mismatch with mm_extractor dim {beit_config.encoder_embed_dim}"
|
||||
out_dim = config.out_dim
|
||||
text_fc = [
|
||||
nn.Linear(in_dim, in_dim),
|
||||
nn.ReLU(),
|
||||
nn.Linear(in_dim, out_dim)
|
||||
]
|
||||
self.text_hidden_fcs = nn.ModuleList([nn.Sequential(*text_fc)])
|
||||
self.text_hidden_fcs.train()
|
||||
for param in self.text_hidden_fcs.parameters():
|
||||
param.requires_grad = True
|
||||
|
||||
def postprocess_masks(self, masks: torch.Tensor, orig_hw) -> torch.Tensor:
|
||||
"""
|
||||
Perform PostProcessing on output masks.
|
||||
"""
|
||||
masks = masks.float()
|
||||
masks = F.interpolate(masks, orig_hw, mode="bilinear", align_corners=False)
|
||||
return masks
|
||||
|
||||
def forward(
|
||||
self,
|
||||
images: torch.FloatTensor,
|
||||
images_evf: torch.FloatTensor,
|
||||
input_ids: torch.LongTensor,
|
||||
attention_masks: torch.LongTensor,
|
||||
offset: torch.LongTensor,
|
||||
masks_list: List[torch.FloatTensor],
|
||||
label_list: List[torch.Tensor],
|
||||
resize_list: List[tuple],
|
||||
inference: bool = False,
|
||||
**kwargs,
|
||||
):
|
||||
# image_embeddings = self.get_visual_embs(images)
|
||||
backbone_out = self.visual_model.forward_image(images)
|
||||
# dict_keys(['vision_features', 'vision_pos_enc', 'backbone_fpn'])
|
||||
_, image_embeddings, _, _ = self.visual_model._prepare_backbone_features(backbone_out)
|
||||
image_embeddings = [_.to(images.dtype) for _ in image_embeddings]
|
||||
batch_size = images.shape[0]
|
||||
if self.visual_model.directly_add_no_mem_embed:
|
||||
image_embeddings[-1] = image_embeddings[-1] + self.visual_model.no_mem_embed
|
||||
|
||||
feats = [
|
||||
feat.permute(1, 2, 0).view(batch_size, -1, *feat_size)
|
||||
for feat, feat_size in zip(image_embeddings[::-1], self._bb_feat_sizes[::-1])
|
||||
][::-1]
|
||||
_features = {"image_embed": feats[-1], "high_res_feats": feats[:-1]}
|
||||
|
||||
|
||||
assert batch_size == len(offset) - 1
|
||||
|
||||
images_evf_list = []
|
||||
for i in range(len(offset) - 1):
|
||||
start_i, end_i = offset[i], offset[i + 1]
|
||||
images_evf_i = (
|
||||
images_evf[i]
|
||||
.unsqueeze(0)
|
||||
.expand(end_i - start_i, -1, -1, -1)
|
||||
.contiguous()
|
||||
)
|
||||
images_evf_list.append(images_evf_i)
|
||||
images_evf = torch.cat(images_evf_list, dim=0)
|
||||
|
||||
multimask_output = False
|
||||
output = self.mm_extractor.beit3(
|
||||
visual_tokens=images_evf,
|
||||
textual_tokens=input_ids,
|
||||
text_padding_position=~attention_masks
|
||||
)
|
||||
|
||||
feat = output["encoder_out"][:, :1, ...]
|
||||
|
||||
feat = self.text_hidden_fcs[0](feat)
|
||||
feat = torch.split(feat, [offset[i+1] - offset[i] for i in range(len(offset)-1)])
|
||||
|
||||
pred_masks = []
|
||||
|
||||
for i in range(len(feat)):
|
||||
(
|
||||
sparse_embeddings,
|
||||
dense_embeddings,
|
||||
) = self.visual_model.sam_prompt_encoder(
|
||||
points=None,
|
||||
boxes=None,
|
||||
masks=None,
|
||||
text_embeds=feat[i],
|
||||
)
|
||||
sparse_embeddings = sparse_embeddings.to(feat[i].dtype)
|
||||
high_res_features = [
|
||||
feat_level[i].unsqueeze(0)
|
||||
for feat_level in _features["high_res_feats"]
|
||||
]
|
||||
low_res_masks, iou_predictions, _, _ = self.visual_model.sam_mask_decoder(
|
||||
image_embeddings=_features["image_embed"][i].unsqueeze(0),
|
||||
image_pe=self.visual_model.sam_prompt_encoder.get_dense_pe(),
|
||||
sparse_prompt_embeddings=sparse_embeddings,
|
||||
dense_prompt_embeddings=dense_embeddings,
|
||||
multimask_output=multimask_output,
|
||||
repeat_image = True,
|
||||
high_res_features=high_res_features,
|
||||
)
|
||||
|
||||
if multimask_output:
|
||||
sorted_ids = torch.argsort(iou_predictions, dim=-1, descending=True)
|
||||
low_res_masks = torch.take_along_dim(low_res_masks, sorted_ids[..., None, None], dim=1)[:, :1]
|
||||
|
||||
pred_mask = self.postprocess_masks(
|
||||
low_res_masks,
|
||||
orig_hw=label_list[i].shape,
|
||||
)
|
||||
pred_masks.append(pred_mask[:, 0])
|
||||
|
||||
gt_masks = masks_list
|
||||
|
||||
if inference:
|
||||
return {
|
||||
"pred_masks": pred_masks,
|
||||
"gt_masks": gt_masks,
|
||||
}
|
||||
|
||||
mask_bce_loss = 0
|
||||
mask_dice_loss = 0
|
||||
num_masks = 0
|
||||
for batch_idx in range(len(pred_masks)):
|
||||
gt_mask = gt_masks[batch_idx]
|
||||
pred_mask = pred_masks[batch_idx]
|
||||
|
||||
assert (
|
||||
gt_mask.shape[0] == pred_mask.shape[0]
|
||||
), "gt_mask.shape: {}, pred_mask.shape: {}".format(
|
||||
gt_mask.shape, pred_mask.shape
|
||||
)
|
||||
mask_bce_loss += (
|
||||
sigmoid_ce_loss(pred_mask, gt_mask, num_masks=gt_mask.shape[0])
|
||||
* gt_mask.shape[0]
|
||||
)
|
||||
mask_dice_loss += (
|
||||
dice_loss(pred_mask, gt_mask, num_masks=gt_mask.shape[0])
|
||||
* gt_mask.shape[0]
|
||||
)
|
||||
num_masks += gt_mask.shape[0]
|
||||
|
||||
mask_bce_loss = self.bce_loss_weight * mask_bce_loss / (num_masks + 1e-8)
|
||||
mask_dice_loss = self.dice_loss_weight * mask_dice_loss / (num_masks + 1e-8)
|
||||
mask_loss = mask_bce_loss + mask_dice_loss
|
||||
|
||||
loss = mask_loss
|
||||
|
||||
return {
|
||||
"loss": loss,
|
||||
"mask_bce_loss": mask_bce_loss,
|
||||
"mask_dice_loss": mask_dice_loss,
|
||||
"mask_loss": mask_loss,
|
||||
}
|
||||
|
||||
def inference(
|
||||
self,
|
||||
images,
|
||||
images_evf,
|
||||
input_ids,
|
||||
resize_list,
|
||||
original_size_list,
|
||||
multimask_output=False,
|
||||
):
|
||||
with torch.no_grad():
|
||||
backbone_out = self.visual_model.forward_image(images)
|
||||
# dict_keys(['vision_features', 'vision_pos_enc', 'backbone_fpn'])
|
||||
_, image_embeddings, _, _ = self.visual_model._prepare_backbone_features(backbone_out)
|
||||
image_embeddings = [_.to(images.dtype) for _ in image_embeddings]
|
||||
batch_size = images.shape[0]
|
||||
if self.visual_model.directly_add_no_mem_embed:
|
||||
image_embeddings[-1] = image_embeddings[-1] + self.visual_model.no_mem_embed
|
||||
|
||||
feats = [
|
||||
feat.permute(1, 2, 0).view(batch_size, -1, *feat_size)
|
||||
for feat, feat_size in zip(image_embeddings[::-1], self._bb_feat_sizes[::-1])
|
||||
][::-1]
|
||||
_features = {"image_embed": feats[-1], "high_res_feats": feats[:-1]}
|
||||
|
||||
|
||||
multimask_output = multimask_output
|
||||
|
||||
output = self.mm_extractor.beit3(visual_tokens=images_evf, textual_tokens=input_ids, text_padding_position=torch.zeros_like(input_ids))
|
||||
|
||||
feat = output["encoder_out"][:, :1, ...]
|
||||
feat = self.text_hidden_fcs[0](feat)
|
||||
(
|
||||
sparse_embeddings,
|
||||
dense_embeddings,
|
||||
) = self.visual_model.sam_prompt_encoder(
|
||||
points=None,
|
||||
boxes=None,
|
||||
masks=None,
|
||||
text_embeds=feat,
|
||||
)
|
||||
high_res_features = _features["high_res_feats"]
|
||||
sparse_embeddings = sparse_embeddings.to(feat.dtype)
|
||||
low_res_masks, iou_predictions, _, _ = self.visual_model.sam_mask_decoder(
|
||||
image_embeddings=_features["image_embed"],
|
||||
image_pe=self.visual_model.sam_prompt_encoder.get_dense_pe(),
|
||||
sparse_prompt_embeddings=sparse_embeddings,
|
||||
dense_prompt_embeddings=dense_embeddings,
|
||||
multimask_output=multimask_output,
|
||||
repeat_image = True,
|
||||
high_res_features=high_res_features,
|
||||
)
|
||||
if multimask_output:
|
||||
sorted_ids = torch.argsort(iou_predictions, dim=-1, descending=True)
|
||||
low_res_masks = torch.take_along_dim(low_res_masks, sorted_ids[..., None, None], dim=1)[:, :1]
|
||||
|
||||
pred_mask = self.postprocess_masks(
|
||||
low_res_masks,
|
||||
orig_hw=original_size_list[0],
|
||||
)
|
||||
|
||||
return pred_mask[:, 0]
|
||||
|
||||
|
||||
AutoConfig.register("evf", EvfConfig)
|
||||
AutoModelForCausalLM.register(EvfConfig, EvfSam2Model)
|
||||
@@ -0,0 +1,321 @@
|
||||
from typing import List
|
||||
import os
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from transformers import PreTrainedModel, AutoConfig, AutoModelForCausalLM
|
||||
from .segment_anything_2.sam2.build_sam import build_sam2, build_sam2_video_predictor
|
||||
from .unilm.beit3.modeling_utils import BEiT3Wrapper, _get_base_config, _get_large_config
|
||||
from .configuration_evf import EvfConfig
|
||||
from .segment_anything_2.sam2.utils.misc import load_video_frames
|
||||
from collections import OrderedDict
|
||||
|
||||
|
||||
|
||||
def dice_loss(
|
||||
inputs: torch.Tensor,
|
||||
targets: torch.Tensor,
|
||||
num_masks: float,
|
||||
scale=1000, # 100000.0,
|
||||
eps=1e-6,
|
||||
):
|
||||
"""
|
||||
Compute the DICE loss, similar to generalized IOU for masks
|
||||
Args:
|
||||
inputs: A float tensor of arbitrary shape.
|
||||
The predictions for each example.
|
||||
targets: A float tensor with the same shape as inputs. Stores the binary
|
||||
classification label for each element in inputs
|
||||
(0 for the negative class and 1 for the positive class).
|
||||
"""
|
||||
inputs = inputs.sigmoid()
|
||||
inputs = inputs.flatten(1, 2)
|
||||
targets = targets.flatten(1, 2)
|
||||
numerator = 2 * (inputs / scale * targets).sum(-1)
|
||||
denominator = (inputs / scale).sum(-1) + (targets / scale).sum(-1)
|
||||
loss = 1 - (numerator + eps) / (denominator + eps)
|
||||
loss = loss.sum() / (num_masks + 1e-8)
|
||||
return loss
|
||||
|
||||
|
||||
def sigmoid_ce_loss(
|
||||
inputs: torch.Tensor,
|
||||
targets: torch.Tensor,
|
||||
num_masks: float,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
inputs: A float tensor of arbitrary shape.
|
||||
The predictions for each example.
|
||||
targets: A float tensor with the same shape as inputs. Stores the binary
|
||||
classification label for each element in inputs
|
||||
(0 for the negative class and 1 for the positive class).
|
||||
Returns:
|
||||
Loss tensor
|
||||
"""
|
||||
loss = F.binary_cross_entropy_with_logits(inputs, targets, reduction="none")
|
||||
loss = loss.flatten(1, 2).mean(1).sum() / (num_masks + 1e-8)
|
||||
return loss
|
||||
|
||||
class EvfSam2Model(PreTrainedModel):
|
||||
config_class = EvfConfig
|
||||
def __init__(
|
||||
self,
|
||||
config,
|
||||
**kwargs
|
||||
):
|
||||
super(EvfSam2Model, self).__init__(config)
|
||||
|
||||
self.config = config
|
||||
self.vision_pretrained = kwargs.get("vision_pretrained", None)
|
||||
self.encoder_pretrained = kwargs.get("encoder_pretrained", None)
|
||||
self.dice_loss_weight = kwargs.get("dice_loss_weight", None)
|
||||
self.bce_loss_weight = kwargs.get("bce_loss_weight", None)
|
||||
self.train_mask_decoder = kwargs.get("train_mask_decoder", False)
|
||||
self.train_prompt_encoder = kwargs.get("train_prompt_encoder", False)
|
||||
self.initialize_evf_modules(config)
|
||||
self._bb_feat_sizes = [
|
||||
(256, 256),
|
||||
(128, 128),
|
||||
(64, 64),
|
||||
]
|
||||
|
||||
def initialize_evf_modules(self, config):
|
||||
# SAM
|
||||
if config.sam_scale=="large":
|
||||
self.visual_model = build_sam2_video_predictor("sam2_hiera_l.yaml", self.vision_pretrained, device=None)
|
||||
elif config.sam_scale=="tiny":
|
||||
self.visual_model = build_sam2_video_predictor("sam2_hiera_t.yaml", self.vision_pretrained, device=None)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
for param in self.visual_model.parameters():
|
||||
param.requires_grad = False
|
||||
if self.train_mask_decoder:
|
||||
self.visual_model.sam_mask_decoder.train()
|
||||
for param in self.visual_model.sam_mask_decoder.parameters():
|
||||
param.requires_grad = True
|
||||
if self.train_prompt_encoder:
|
||||
self.visual_model.sam_prompt_encoder.no_mask_embed.requires_grad_(True)
|
||||
|
||||
# beit-3
|
||||
if self.config.mm_extractor_scale == "base":
|
||||
beit_config = _get_base_config()
|
||||
elif self.config.mm_extractor_scale == "large":
|
||||
beit_config = _get_large_config()
|
||||
else:
|
||||
raise AttributeError(f"model config should contain key 'mm_extractor_scale', with value 'base' or 'large'.")
|
||||
|
||||
self.mm_extractor = BEiT3Wrapper(beit_config)
|
||||
if self.encoder_pretrained is not None:
|
||||
beit_state_dict = torch.load(self.encoder_pretrained)["model"]
|
||||
self.mm_extractor.load_state_dict(
|
||||
beit_state_dict,
|
||||
strict=False
|
||||
)
|
||||
|
||||
for param in self.mm_extractor.parameters():
|
||||
param.requires_grad = True
|
||||
|
||||
# Projection layer
|
||||
in_dim = config.hidden_size
|
||||
assert in_dim==beit_config.encoder_embed_dim, \
|
||||
f"projection layer dim {in_dim} mismatch with mm_extractor dim {beit_config.encoder_embed_dim}"
|
||||
out_dim = config.out_dim
|
||||
text_fc = [
|
||||
nn.Linear(in_dim, in_dim),
|
||||
nn.ReLU(),
|
||||
nn.Linear(in_dim, out_dim)
|
||||
]
|
||||
self.text_hidden_fcs = nn.ModuleList([nn.Sequential(*text_fc)])
|
||||
self.text_hidden_fcs.train()
|
||||
for param in self.text_hidden_fcs.parameters():
|
||||
param.requires_grad = True
|
||||
|
||||
|
||||
def postprocess_masks(self, masks: torch.Tensor, orig_hw) -> torch.Tensor:
|
||||
"""
|
||||
Perform PostProcessing on output masks.
|
||||
"""
|
||||
masks = masks.float()
|
||||
masks = F.interpolate(masks, orig_hw, mode="bilinear", align_corners=False)
|
||||
return masks
|
||||
|
||||
# def forward(
|
||||
# self,
|
||||
# images: torch.FloatTensor,
|
||||
# images_evf: torch.FloatTensor,
|
||||
# input_ids: torch.LongTensor,
|
||||
# attention_masks: torch.LongTensor,
|
||||
# offset: torch.LongTensor,
|
||||
# masks_list: List[torch.FloatTensor],
|
||||
# label_list: List[torch.Tensor],
|
||||
# resize_list: List[tuple],
|
||||
# inference: bool = False,
|
||||
# **kwargs,
|
||||
# ):
|
||||
# # image_embeddings = self.get_visual_embs(images)
|
||||
# backbone_out = self.visual_model.forward_image(images)
|
||||
# # dict_keys(['vision_features', 'vision_pos_enc', 'backbone_fpn'])
|
||||
# _, image_embeddings, _, _ = self.visual_model._prepare_backbone_features(backbone_out)
|
||||
# image_embeddings = [_.to(images.dtype) for _ in image_embeddings]
|
||||
# batch_size = images.shape[0]
|
||||
# if self.visual_model.directly_add_no_mem_embed:
|
||||
# image_embeddings[-1] = image_embeddings[-1] + self.visual_model.no_mem_embed
|
||||
|
||||
# feats = [
|
||||
# feat.permute(1, 2, 0).view(batch_size, -1, *feat_size)
|
||||
# for feat, feat_size in zip(image_embeddings[::-1], self._bb_feat_sizes[::-1])
|
||||
# ][::-1]
|
||||
# _features = {"image_embed": feats[-1], "high_res_feats": feats[:-1]}
|
||||
|
||||
|
||||
# assert batch_size == len(offset) - 1
|
||||
|
||||
# images_evf_list = []
|
||||
# for i in range(len(offset) - 1):
|
||||
# start_i, end_i = offset[i], offset[i + 1]
|
||||
# images_evf_i = (
|
||||
# images_evf[i]
|
||||
# .unsqueeze(0)
|
||||
# .expand(end_i - start_i, -1, -1, -1)
|
||||
# .contiguous()
|
||||
# )
|
||||
# images_evf_list.append(images_evf_i)
|
||||
# images_evf = torch.cat(images_evf_list, dim=0)
|
||||
|
||||
# multimask_output = False
|
||||
# output = self.mm_extractor.beit3(
|
||||
# visual_tokens=images_evf,
|
||||
# textual_tokens=input_ids,
|
||||
# text_padding_position=~attention_masks
|
||||
# )
|
||||
|
||||
# feat = output["encoder_out"][:, :1, ...]
|
||||
|
||||
# feat = self.text_hidden_fcs[0](feat)
|
||||
# feat = torch.split(feat, [offset[i+1] - offset[i] for i in range(len(offset)-1)])
|
||||
|
||||
# pred_masks = []
|
||||
|
||||
# for i in range(len(feat)):
|
||||
# (
|
||||
# sparse_embeddings,
|
||||
# dense_embeddings,
|
||||
# ) = self.visual_model.sam_prompt_encoder(
|
||||
# points=None,
|
||||
# boxes=None,
|
||||
# masks=None,
|
||||
# text_embeds=feat[i],
|
||||
# )
|
||||
# sparse_embeddings = sparse_embeddings.to(feat[i].dtype)
|
||||
# high_res_features = [
|
||||
# feat_level[i].unsqueeze(0)
|
||||
# for feat_level in _features["high_res_feats"]
|
||||
# ]
|
||||
# low_res_masks, iou_predictions, _, _ = self.visual_model.sam_mask_decoder(
|
||||
# image_embeddings=_features["image_embed"][i].unsqueeze(0),
|
||||
# image_pe=self.visual_model.sam_prompt_encoder.get_dense_pe(),
|
||||
# sparse_prompt_embeddings=sparse_embeddings,
|
||||
# dense_prompt_embeddings=dense_embeddings,
|
||||
# multimask_output=multimask_output,
|
||||
# repeat_image = True,
|
||||
# high_res_features=high_res_features,
|
||||
# )
|
||||
|
||||
# if multimask_output:
|
||||
# sorted_ids = torch.argsort(iou_predictions, dim=-1, descending=True)
|
||||
# low_res_masks = torch.take_along_dim(low_res_masks, sorted_ids[..., None, None], dim=1)[:, :1]
|
||||
|
||||
# pred_mask = self.postprocess_masks(
|
||||
# low_res_masks,
|
||||
# orig_hw=label_list[i].shape,
|
||||
# )
|
||||
# pred_masks.append(pred_mask[:, 0])
|
||||
|
||||
# gt_masks = masks_list
|
||||
|
||||
# if inference:
|
||||
# return {
|
||||
# "pred_masks": pred_masks,
|
||||
# "gt_masks": gt_masks,
|
||||
# }
|
||||
|
||||
# mask_bce_loss = 0
|
||||
# mask_dice_loss = 0
|
||||
# num_masks = 0
|
||||
# for batch_idx in range(len(pred_masks)):
|
||||
# gt_mask = gt_masks[batch_idx]
|
||||
# pred_mask = pred_masks[batch_idx]
|
||||
|
||||
# assert (
|
||||
# gt_mask.shape[0] == pred_mask.shape[0]
|
||||
# ), "gt_mask.shape: {}, pred_mask.shape: {}".format(
|
||||
# gt_mask.shape, pred_mask.shape
|
||||
# )
|
||||
# mask_bce_loss += (
|
||||
# sigmoid_ce_loss(pred_mask, gt_mask, num_masks=gt_mask.shape[0])
|
||||
# * gt_mask.shape[0]
|
||||
# )
|
||||
# mask_dice_loss += (
|
||||
# dice_loss(pred_mask, gt_mask, num_masks=gt_mask.shape[0])
|
||||
# * gt_mask.shape[0]
|
||||
# )
|
||||
# num_masks += gt_mask.shape[0]
|
||||
|
||||
# mask_bce_loss = self.bce_loss_weight * mask_bce_loss / (num_masks + 1e-8)
|
||||
# mask_dice_loss = self.dice_loss_weight * mask_dice_loss / (num_masks + 1e-8)
|
||||
# mask_loss = mask_bce_loss + mask_dice_loss
|
||||
|
||||
# loss = mask_loss
|
||||
|
||||
# return {
|
||||
# "loss": loss,
|
||||
# "mask_bce_loss": mask_bce_loss,
|
||||
# "mask_dice_loss": mask_dice_loss,
|
||||
# "mask_loss": mask_loss,
|
||||
# }
|
||||
|
||||
def inference(
|
||||
self,
|
||||
video_path,
|
||||
images_evf,
|
||||
input_ids,
|
||||
# original_size_list,
|
||||
multimask_output=False,
|
||||
):
|
||||
predictor = self.visual_model
|
||||
inference_state = predictor.init_state(video_path=video_path)
|
||||
predictor.reset_state(inference_state)
|
||||
|
||||
|
||||
multimask_output = multimask_output
|
||||
|
||||
output = self.mm_extractor.beit3(visual_tokens=images_evf, textual_tokens=input_ids, text_padding_position=torch.zeros_like(input_ids))
|
||||
|
||||
feat = output["encoder_out"][:, :1, ...]
|
||||
feat = self.text_hidden_fcs[0](feat)
|
||||
|
||||
ann_frame_idx = 0 # the frame index we interact with
|
||||
ann_obj_id = 1 # give a unique id to each object we interact with (it can be any integers)
|
||||
|
||||
_, out_obj_ids, out_mask_logits = predictor.add_new_text(
|
||||
inference_state=inference_state,
|
||||
frame_idx=ann_frame_idx,
|
||||
obj_id=ann_obj_id,
|
||||
text=feat
|
||||
)
|
||||
|
||||
# run propagation throughout the video and collect the results in a dict
|
||||
video_segments = {} # video_segments contains the per-frame segmentation results
|
||||
for out_frame_idx, out_obj_ids, out_mask_logits in predictor.propagate_in_video(inference_state):
|
||||
video_segments[out_frame_idx] = {
|
||||
out_obj_id: (out_mask_logits[i] > 0.0).cpu().numpy()
|
||||
for i, out_obj_id in enumerate(out_obj_ids)
|
||||
}
|
||||
|
||||
return video_segments
|
||||
|
||||
|
||||
AutoConfig.register("evf", EvfConfig)
|
||||
AutoModelForCausalLM.register(EvfConfig, EvfSam2Model)
|
||||
@@ -0,0 +1,10 @@
|
||||
# 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.
|
||||
|
||||
from .automatic_mask_generator import SamAutomaticMaskGenerator
|
||||
from .build_sam import (build_sam, build_sam_vit_b, build_sam_vit_h,
|
||||
build_sam_vit_l, sam_model_registry)
|
||||
from .predictor import SamPredictor
|
||||
@@ -0,0 +1,372 @@
|
||||
# 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.
|
||||
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from torchvision.ops.boxes import batched_nms, box_area # type: ignore
|
||||
|
||||
from .modeling import Sam
|
||||
from .predictor import SamPredictor
|
||||
from .utils.amg import (MaskData, area_from_rle, batch_iterator,
|
||||
batched_mask_to_box, box_xyxy_to_xywh,
|
||||
build_all_layer_point_grids, calculate_stability_score,
|
||||
coco_encode_rle, generate_crop_boxes,
|
||||
is_box_near_crop_edge, mask_to_rle_pytorch,
|
||||
remove_small_regions, rle_to_mask, uncrop_boxes_xyxy,
|
||||
uncrop_masks, uncrop_points)
|
||||
|
||||
|
||||
class SamAutomaticMaskGenerator:
|
||||
def __init__(
|
||||
self,
|
||||
model: Sam,
|
||||
points_per_side: Optional[int] = 32,
|
||||
points_per_batch: int = 64,
|
||||
pred_iou_thresh: float = 0.88,
|
||||
stability_score_thresh: float = 0.95,
|
||||
stability_score_offset: float = 1.0,
|
||||
box_nms_thresh: float = 0.7,
|
||||
crop_n_layers: int = 0,
|
||||
crop_nms_thresh: float = 0.7,
|
||||
crop_overlap_ratio: float = 512 / 1500,
|
||||
crop_n_points_downscale_factor: int = 1,
|
||||
point_grids: Optional[List[np.ndarray]] = None,
|
||||
min_mask_region_area: int = 0,
|
||||
output_mode: str = "binary_mask",
|
||||
) -> None:
|
||||
"""
|
||||
Using a SAM model, generates masks for the entire image.
|
||||
Generates a grid of point prompts over the image, then filters
|
||||
low quality and duplicate masks. The default settings are chosen
|
||||
for SAM with a ViT-H backbone.
|
||||
|
||||
Arguments:
|
||||
model (Sam): The SAM model to use for mask prediction.
|
||||
points_per_side (int or None): The number of points to be sampled
|
||||
along one side of the image. The total number of points is
|
||||
points_per_side**2. If None, 'point_grids' must provide explicit
|
||||
point sampling.
|
||||
points_per_batch (int): Sets the number of points run simultaneously
|
||||
by the model. Higher numbers may be faster but use more GPU memory.
|
||||
pred_iou_thresh (float): A filtering threshold in [0,1], using the
|
||||
model's predicted mask quality.
|
||||
stability_score_thresh (float): A filtering threshold in [0,1], using
|
||||
the stability of the mask under changes to the cutoff used to binarize
|
||||
the model's mask predictions.
|
||||
stability_score_offset (float): The amount to shift the cutoff when
|
||||
calculated the stability score.
|
||||
box_nms_thresh (float): The box IoU cutoff used by non-maximal
|
||||
suppression to filter duplicate masks.
|
||||
crop_n_layers (int): If >0, mask prediction will be run again on
|
||||
crops of the image. Sets the number of layers to run, where each
|
||||
layer has 2**i_layer number of image crops.
|
||||
crop_nms_thresh (float): The box IoU cutoff used by non-maximal
|
||||
suppression to filter duplicate masks between different crops.
|
||||
crop_overlap_ratio (float): Sets the degree to which crops overlap.
|
||||
In the first crop layer, crops will overlap by this fraction of
|
||||
the image length. Later layers with more crops scale down this overlap.
|
||||
crop_n_points_downscale_factor (int): The number of points-per-side
|
||||
sampled in layer n is scaled down by crop_n_points_downscale_factor**n.
|
||||
point_grids (list(np.ndarray) or None): A list over explicit grids
|
||||
of points used for sampling, normalized to [0,1]. The nth grid in the
|
||||
list is used in the nth crop layer. Exclusive with points_per_side.
|
||||
min_mask_region_area (int): If >0, postprocessing will be applied
|
||||
to remove disconnected regions and holes in masks with area smaller
|
||||
than min_mask_region_area. Requires opencv.
|
||||
output_mode (str): The form masks are returned in. Can be 'binary_mask',
|
||||
'uncompressed_rle', or 'coco_rle'. 'coco_rle' requires pycocotools.
|
||||
For large resolutions, 'binary_mask' may consume large amounts of
|
||||
memory.
|
||||
"""
|
||||
|
||||
assert (points_per_side is None) != (
|
||||
point_grids is None
|
||||
), "Exactly one of points_per_side or point_grid must be provided."
|
||||
if points_per_side is not None:
|
||||
self.point_grids = build_all_layer_point_grids(
|
||||
points_per_side,
|
||||
crop_n_layers,
|
||||
crop_n_points_downscale_factor,
|
||||
)
|
||||
elif point_grids is not None:
|
||||
self.point_grids = point_grids
|
||||
else:
|
||||
raise ValueError("Can't have both points_per_side and point_grid be None.")
|
||||
|
||||
assert output_mode in [
|
||||
"binary_mask",
|
||||
"uncompressed_rle",
|
||||
"coco_rle",
|
||||
], f"Unknown output_mode {output_mode}."
|
||||
if output_mode == "coco_rle":
|
||||
from pycocotools import \
|
||||
mask as mask_utils # type: ignore # noqa: F401
|
||||
|
||||
if min_mask_region_area > 0:
|
||||
import cv2 # type: ignore # noqa: F401
|
||||
|
||||
self.predictor = SamPredictor(model)
|
||||
self.points_per_batch = points_per_batch
|
||||
self.pred_iou_thresh = pred_iou_thresh
|
||||
self.stability_score_thresh = stability_score_thresh
|
||||
self.stability_score_offset = stability_score_offset
|
||||
self.box_nms_thresh = box_nms_thresh
|
||||
self.crop_n_layers = crop_n_layers
|
||||
self.crop_nms_thresh = crop_nms_thresh
|
||||
self.crop_overlap_ratio = crop_overlap_ratio
|
||||
self.crop_n_points_downscale_factor = crop_n_points_downscale_factor
|
||||
self.min_mask_region_area = min_mask_region_area
|
||||
self.output_mode = output_mode
|
||||
|
||||
@torch.no_grad()
|
||||
def generate(self, image: np.ndarray) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Generates masks for the given image.
|
||||
|
||||
Arguments:
|
||||
image (np.ndarray): The image to generate masks for, in HWC uint8 format.
|
||||
|
||||
Returns:
|
||||
list(dict(str, any)): A list over records for masks. Each record is
|
||||
a dict containing the following keys:
|
||||
segmentation (dict(str, any) or np.ndarray): The mask. If
|
||||
output_mode='binary_mask', is an array of shape HW. Otherwise,
|
||||
is a dictionary containing the RLE.
|
||||
bbox (list(float)): The box around the mask, in XYWH format.
|
||||
area (int): The area in pixels of the mask.
|
||||
predicted_iou (float): The model's own prediction of the mask's
|
||||
quality. This is filtered by the pred_iou_thresh parameter.
|
||||
point_coords (list(list(float))): The point coordinates input
|
||||
to the model to generate this mask.
|
||||
stability_score (float): A measure of the mask's quality. This
|
||||
is filtered on using the stability_score_thresh parameter.
|
||||
crop_box (list(float)): The crop of the image used to generate
|
||||
the mask, given in XYWH format.
|
||||
"""
|
||||
|
||||
# Generate masks
|
||||
mask_data = self._generate_masks(image)
|
||||
|
||||
# Filter small disconnected regions and holes in masks
|
||||
if self.min_mask_region_area > 0:
|
||||
mask_data = self.postprocess_small_regions(
|
||||
mask_data,
|
||||
self.min_mask_region_area,
|
||||
max(self.box_nms_thresh, self.crop_nms_thresh),
|
||||
)
|
||||
|
||||
# Encode masks
|
||||
if self.output_mode == "coco_rle":
|
||||
mask_data["segmentations"] = [
|
||||
coco_encode_rle(rle) for rle in mask_data["rles"]
|
||||
]
|
||||
elif self.output_mode == "binary_mask":
|
||||
mask_data["segmentations"] = [rle_to_mask(rle) for rle in mask_data["rles"]]
|
||||
else:
|
||||
mask_data["segmentations"] = mask_data["rles"]
|
||||
|
||||
# Write mask records
|
||||
curr_anns = []
|
||||
for idx in range(len(mask_data["segmentations"])):
|
||||
ann = {
|
||||
"segmentation": mask_data["segmentations"][idx],
|
||||
"area": area_from_rle(mask_data["rles"][idx]),
|
||||
"bbox": box_xyxy_to_xywh(mask_data["boxes"][idx]).tolist(),
|
||||
"predicted_iou": mask_data["iou_preds"][idx].item(),
|
||||
"point_coords": [mask_data["points"][idx].tolist()],
|
||||
"stability_score": mask_data["stability_score"][idx].item(),
|
||||
"crop_box": box_xyxy_to_xywh(mask_data["crop_boxes"][idx]).tolist(),
|
||||
}
|
||||
curr_anns.append(ann)
|
||||
|
||||
return curr_anns
|
||||
|
||||
def _generate_masks(self, image: np.ndarray) -> MaskData:
|
||||
orig_size = image.shape[:2]
|
||||
crop_boxes, layer_idxs = generate_crop_boxes(
|
||||
orig_size, self.crop_n_layers, self.crop_overlap_ratio
|
||||
)
|
||||
|
||||
# Iterate over image crops
|
||||
data = MaskData()
|
||||
for crop_box, layer_idx in zip(crop_boxes, layer_idxs):
|
||||
crop_data = self._process_crop(image, crop_box, layer_idx, orig_size)
|
||||
data.cat(crop_data)
|
||||
|
||||
# Remove duplicate masks between crops
|
||||
if len(crop_boxes) > 1:
|
||||
# Prefer masks from smaller crops
|
||||
scores = 1 / box_area(data["crop_boxes"])
|
||||
scores = scores.to(data["boxes"].device)
|
||||
keep_by_nms = batched_nms(
|
||||
data["boxes"].float(),
|
||||
scores,
|
||||
torch.zeros_like(data["boxes"][:, 0]), # categories
|
||||
iou_threshold=self.crop_nms_thresh,
|
||||
)
|
||||
data.filter(keep_by_nms)
|
||||
|
||||
data.to_numpy()
|
||||
return data
|
||||
|
||||
def _process_crop(
|
||||
self,
|
||||
image: np.ndarray,
|
||||
crop_box: List[int],
|
||||
crop_layer_idx: int,
|
||||
orig_size: Tuple[int, ...],
|
||||
) -> MaskData:
|
||||
# Crop the image and calculate embeddings
|
||||
x0, y0, x1, y1 = crop_box
|
||||
cropped_im = image[y0:y1, x0:x1, :]
|
||||
cropped_im_size = cropped_im.shape[:2]
|
||||
self.predictor.set_image(cropped_im)
|
||||
|
||||
# Get points for this crop
|
||||
points_scale = np.array(cropped_im_size)[None, ::-1]
|
||||
points_for_image = self.point_grids[crop_layer_idx] * points_scale
|
||||
|
||||
# Generate masks for this crop in batches
|
||||
data = MaskData()
|
||||
for (points,) in batch_iterator(self.points_per_batch, points_for_image):
|
||||
batch_data = self._process_batch(
|
||||
points, cropped_im_size, crop_box, orig_size
|
||||
)
|
||||
data.cat(batch_data)
|
||||
del batch_data
|
||||
self.predictor.reset_image()
|
||||
|
||||
# Remove duplicates within this crop.
|
||||
keep_by_nms = batched_nms(
|
||||
data["boxes"].float(),
|
||||
data["iou_preds"],
|
||||
torch.zeros_like(data["boxes"][:, 0]), # categories
|
||||
iou_threshold=self.box_nms_thresh,
|
||||
)
|
||||
data.filter(keep_by_nms)
|
||||
|
||||
# Return to the original image frame
|
||||
data["boxes"] = uncrop_boxes_xyxy(data["boxes"], crop_box)
|
||||
data["points"] = uncrop_points(data["points"], crop_box)
|
||||
data["crop_boxes"] = torch.tensor([crop_box for _ in range(len(data["rles"]))])
|
||||
|
||||
return data
|
||||
|
||||
def _process_batch(
|
||||
self,
|
||||
points: np.ndarray,
|
||||
im_size: Tuple[int, ...],
|
||||
crop_box: List[int],
|
||||
orig_size: Tuple[int, ...],
|
||||
) -> MaskData:
|
||||
orig_h, orig_w = orig_size
|
||||
|
||||
# Run model on this batch
|
||||
transformed_points = self.predictor.transform.apply_coords(points, im_size)
|
||||
in_points = torch.as_tensor(transformed_points, device=self.predictor.device)
|
||||
in_labels = torch.ones(
|
||||
in_points.shape[0], dtype=torch.int, device=in_points.device
|
||||
)
|
||||
masks, iou_preds, _ = self.predictor.predict_torch(
|
||||
in_points[:, None, :],
|
||||
in_labels[:, None],
|
||||
multimask_output=True,
|
||||
return_logits=True,
|
||||
)
|
||||
|
||||
# Serialize predictions and store in MaskData
|
||||
data = MaskData(
|
||||
masks=masks.flatten(0, 1),
|
||||
iou_preds=iou_preds.flatten(0, 1),
|
||||
points=torch.as_tensor(points.repeat(masks.shape[1], axis=0)),
|
||||
)
|
||||
del masks
|
||||
|
||||
# Filter by predicted IoU
|
||||
if self.pred_iou_thresh > 0.0:
|
||||
keep_mask = data["iou_preds"] > self.pred_iou_thresh
|
||||
data.filter(keep_mask)
|
||||
|
||||
# Calculate stability score
|
||||
data["stability_score"] = calculate_stability_score(
|
||||
data["masks"],
|
||||
self.predictor.model.mask_threshold,
|
||||
self.stability_score_offset,
|
||||
)
|
||||
if self.stability_score_thresh > 0.0:
|
||||
keep_mask = data["stability_score"] >= self.stability_score_thresh
|
||||
data.filter(keep_mask)
|
||||
|
||||
# Threshold masks and calculate boxes
|
||||
data["masks"] = data["masks"] > self.predictor.model.mask_threshold
|
||||
data["boxes"] = batched_mask_to_box(data["masks"])
|
||||
|
||||
# Filter boxes that touch crop boundaries
|
||||
keep_mask = ~is_box_near_crop_edge(
|
||||
data["boxes"], crop_box, [0, 0, orig_w, orig_h]
|
||||
)
|
||||
if not torch.all(keep_mask):
|
||||
data.filter(keep_mask)
|
||||
|
||||
# Compress to RLE
|
||||
data["masks"] = uncrop_masks(data["masks"], crop_box, orig_h, orig_w)
|
||||
data["rles"] = mask_to_rle_pytorch(data["masks"])
|
||||
del data["masks"]
|
||||
|
||||
return data
|
||||
|
||||
@staticmethod
|
||||
def postprocess_small_regions(
|
||||
mask_data: MaskData, min_area: int, nms_thresh: float
|
||||
) -> MaskData:
|
||||
"""
|
||||
Removes small disconnected regions and holes in masks, then reruns
|
||||
box NMS to remove any new duplicates.
|
||||
|
||||
Edits mask_data in place.
|
||||
|
||||
Requires open-cv as a dependency.
|
||||
"""
|
||||
if len(mask_data["rles"]) == 0:
|
||||
return mask_data
|
||||
|
||||
# Filter small disconnected regions and holes
|
||||
new_masks = []
|
||||
scores = []
|
||||
for rle in mask_data["rles"]:
|
||||
mask = rle_to_mask(rle)
|
||||
|
||||
mask, changed = remove_small_regions(mask, min_area, mode="holes")
|
||||
unchanged = not changed
|
||||
mask, changed = remove_small_regions(mask, min_area, mode="islands")
|
||||
unchanged = unchanged and not changed
|
||||
|
||||
new_masks.append(torch.as_tensor(mask).unsqueeze(0))
|
||||
# Give score=0 to changed masks and score=1 to unchanged masks
|
||||
# so NMS will prefer ones that didn't need postprocessing
|
||||
scores.append(float(unchanged))
|
||||
|
||||
# Recalculate boxes and remove any new duplicates
|
||||
masks = torch.cat(new_masks, dim=0)
|
||||
boxes = batched_mask_to_box(masks)
|
||||
keep_by_nms = batched_nms(
|
||||
boxes.float(),
|
||||
torch.as_tensor(scores),
|
||||
torch.zeros_like(boxes[:, 0]), # categories
|
||||
iou_threshold=nms_thresh,
|
||||
)
|
||||
|
||||
# Only recalculate RLEs for masks that have changed
|
||||
for i_mask in keep_by_nms:
|
||||
if scores[i_mask] == 0.0:
|
||||
mask_torch = masks[i_mask].unsqueeze(0)
|
||||
mask_data["rles"][i_mask] = mask_to_rle_pytorch(mask_torch)[0]
|
||||
mask_data["boxes"][i_mask] = boxes[i_mask] # update res directly
|
||||
mask_data.filter(keep_by_nms)
|
||||
|
||||
return mask_data
|
||||
@@ -0,0 +1,108 @@
|
||||
# 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.
|
||||
|
||||
from functools import partial
|
||||
|
||||
import torch
|
||||
|
||||
from .modeling import (ImageEncoderViT, MaskDecoder, PromptEncoder, Sam,
|
||||
TwoWayTransformer)
|
||||
|
||||
|
||||
def build_sam_vit_h(checkpoint=None):
|
||||
return _build_sam(
|
||||
encoder_embed_dim=1280,
|
||||
encoder_depth=32,
|
||||
encoder_num_heads=16,
|
||||
encoder_global_attn_indexes=[7, 15, 23, 31],
|
||||
checkpoint=checkpoint,
|
||||
)
|
||||
|
||||
|
||||
build_sam = build_sam_vit_h
|
||||
|
||||
|
||||
def build_sam_vit_l(checkpoint=None):
|
||||
return _build_sam(
|
||||
encoder_embed_dim=1024,
|
||||
encoder_depth=24,
|
||||
encoder_num_heads=16,
|
||||
encoder_global_attn_indexes=[5, 11, 17, 23],
|
||||
checkpoint=checkpoint,
|
||||
)
|
||||
|
||||
|
||||
def build_sam_vit_b(checkpoint=None):
|
||||
return _build_sam(
|
||||
encoder_embed_dim=768,
|
||||
encoder_depth=12,
|
||||
encoder_num_heads=12,
|
||||
encoder_global_attn_indexes=[2, 5, 8, 11],
|
||||
checkpoint=checkpoint,
|
||||
)
|
||||
|
||||
|
||||
sam_model_registry = {
|
||||
"default": build_sam_vit_h,
|
||||
"vit_h": build_sam_vit_h,
|
||||
"vit_l": build_sam_vit_l,
|
||||
"vit_b": build_sam_vit_b,
|
||||
}
|
||||
|
||||
|
||||
def _build_sam(
|
||||
encoder_embed_dim,
|
||||
encoder_depth,
|
||||
encoder_num_heads,
|
||||
encoder_global_attn_indexes,
|
||||
checkpoint=None,
|
||||
):
|
||||
prompt_embed_dim = 256
|
||||
image_size = 1024
|
||||
vit_patch_size = 16
|
||||
image_embedding_size = image_size // vit_patch_size
|
||||
sam = Sam(
|
||||
image_encoder=ImageEncoderViT(
|
||||
depth=encoder_depth,
|
||||
embed_dim=encoder_embed_dim,
|
||||
img_size=image_size,
|
||||
mlp_ratio=4,
|
||||
norm_layer=partial(torch.nn.LayerNorm, eps=1e-6),
|
||||
num_heads=encoder_num_heads,
|
||||
patch_size=vit_patch_size,
|
||||
qkv_bias=True,
|
||||
use_rel_pos=True,
|
||||
global_attn_indexes=encoder_global_attn_indexes,
|
||||
window_size=14,
|
||||
out_chans=prompt_embed_dim,
|
||||
),
|
||||
prompt_encoder=PromptEncoder(
|
||||
embed_dim=prompt_embed_dim,
|
||||
image_embedding_size=(image_embedding_size, image_embedding_size),
|
||||
input_image_size=(image_size, image_size),
|
||||
mask_in_chans=16,
|
||||
),
|
||||
mask_decoder=MaskDecoder(
|
||||
num_multimask_outputs=3,
|
||||
transformer=TwoWayTransformer(
|
||||
depth=2,
|
||||
embedding_dim=prompt_embed_dim,
|
||||
mlp_dim=2048,
|
||||
num_heads=8,
|
||||
),
|
||||
transformer_dim=prompt_embed_dim,
|
||||
iou_head_depth=3,
|
||||
iou_head_hidden_dim=256,
|
||||
),
|
||||
pixel_mean=[123.675, 116.28, 103.53],
|
||||
pixel_std=[58.395, 57.12, 57.375],
|
||||
)
|
||||
sam.eval()
|
||||
if checkpoint is not None:
|
||||
with open(checkpoint, "rb") as f:
|
||||
state_dict = torch.load(f)
|
||||
sam.load_state_dict(state_dict, strict=False)
|
||||
return sam
|
||||
@@ -0,0 +1,11 @@
|
||||
# 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.
|
||||
|
||||
from .image_encoder import ImageEncoderViT
|
||||
from .mask_decoder import MaskDecoder
|
||||
from .prompt_encoder import PromptEncoder
|
||||
from .sam import Sam
|
||||
from .transformer import TwoWayTransformer
|
||||
@@ -0,0 +1,43 @@
|
||||
# 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.
|
||||
|
||||
from typing import Type
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class MLPBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
embedding_dim: int,
|
||||
mlp_dim: int,
|
||||
act: Type[nn.Module] = nn.GELU,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.lin1 = nn.Linear(embedding_dim, mlp_dim)
|
||||
self.lin2 = nn.Linear(mlp_dim, embedding_dim)
|
||||
self.act = act()
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return self.lin2(self.act(self.lin1(x)))
|
||||
|
||||
|
||||
# From https://github.com/facebookresearch/detectron2/blob/main/detectron2/layers/batch_norm.py # noqa
|
||||
# Itself from https://github.com/facebookresearch/ConvNeXt/blob/d1fa8f6fef0a165b27399986cc2bdacc92777e40/models/convnext.py#L119 # noqa
|
||||
class LayerNorm2d(nn.Module):
|
||||
def __init__(self, num_channels: int, eps: float = 1e-6) -> None:
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.ones(num_channels))
|
||||
self.bias = nn.Parameter(torch.zeros(num_channels))
|
||||
self.eps = eps
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
u = x.mean(1, keepdim=True)
|
||||
s = (x - u).pow(2).mean(1, keepdim=True)
|
||||
x = (x - u) / torch.sqrt(s + self.eps)
|
||||
x = self.weight[:, None, None] * x + self.bias[:, None, None]
|
||||
return x
|
||||
@@ -0,0 +1,426 @@
|
||||
# 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.
|
||||
|
||||
from typing import Optional, Tuple, Type
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from .common import LayerNorm2d, MLPBlock
|
||||
|
||||
|
||||
# This class and its supporting functions below lightly adapted from the ViTDet backbone available at: https://github.com/facebookresearch/detectron2/blob/main/detectron2/modeling/backbone/vit.py # noqa
|
||||
class ImageEncoderViT(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
img_size: int = 1024,
|
||||
patch_size: int = 16,
|
||||
in_chans: int = 3,
|
||||
embed_dim: int = 768,
|
||||
depth: int = 12,
|
||||
num_heads: int = 12,
|
||||
mlp_ratio: float = 4.0,
|
||||
out_chans: int = 256,
|
||||
qkv_bias: bool = True,
|
||||
norm_layer: Type[nn.Module] = nn.LayerNorm,
|
||||
act_layer: Type[nn.Module] = nn.GELU,
|
||||
use_abs_pos: bool = True,
|
||||
use_rel_pos: bool = False,
|
||||
rel_pos_zero_init: bool = True,
|
||||
window_size: int = 0,
|
||||
global_attn_indexes: Tuple[int, ...] = (),
|
||||
) -> None:
|
||||
"""
|
||||
Args:
|
||||
img_size (int): Input image size.
|
||||
patch_size (int): Patch size.
|
||||
in_chans (int): Number of input image channels.
|
||||
embed_dim (int): Patch embedding dimension.
|
||||
depth (int): Depth of ViT.
|
||||
num_heads (int): Number of attention heads in each ViT block.
|
||||
mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.
|
||||
qkv_bias (bool): If True, add a learnable bias to query, key, value.
|
||||
norm_layer (nn.Module): Normalization layer.
|
||||
act_layer (nn.Module): Activation layer.
|
||||
use_abs_pos (bool): If True, use absolute positional embeddings.
|
||||
use_rel_pos (bool): If True, add relative positional embeddings to the attention map.
|
||||
rel_pos_zero_init (bool): If True, zero initialize relative positional parameters.
|
||||
window_size (int): Window size for window attention blocks.
|
||||
global_attn_indexes (list): Indexes for blocks using global attention.
|
||||
"""
|
||||
super().__init__()
|
||||
self.img_size = img_size
|
||||
self.embed_dim = embed_dim
|
||||
self.out_chans = out_chans
|
||||
|
||||
self.patch_embed = PatchEmbed(
|
||||
kernel_size=(patch_size, patch_size),
|
||||
stride=(patch_size, patch_size),
|
||||
in_chans=in_chans,
|
||||
embed_dim=embed_dim,
|
||||
)
|
||||
|
||||
self.pos_embed: Optional[nn.Parameter] = None
|
||||
if use_abs_pos:
|
||||
# Initialize absolute positional embedding with pretrain image size.
|
||||
self.pos_embed = nn.Parameter(
|
||||
torch.zeros(
|
||||
1, img_size // patch_size, img_size // patch_size, embed_dim
|
||||
)
|
||||
)
|
||||
|
||||
self.blocks = nn.ModuleList()
|
||||
for i in range(depth):
|
||||
block = Block(
|
||||
dim=embed_dim,
|
||||
num_heads=num_heads,
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
norm_layer=norm_layer,
|
||||
act_layer=act_layer,
|
||||
use_rel_pos=use_rel_pos,
|
||||
rel_pos_zero_init=rel_pos_zero_init,
|
||||
window_size=window_size if i not in global_attn_indexes else 0,
|
||||
input_size=(img_size // patch_size, img_size // patch_size),
|
||||
)
|
||||
self.blocks.append(block)
|
||||
|
||||
self.neck = nn.Sequential(
|
||||
nn.Conv2d(
|
||||
embed_dim,
|
||||
out_chans,
|
||||
kernel_size=1,
|
||||
bias=False,
|
||||
),
|
||||
LayerNorm2d(out_chans),
|
||||
nn.Conv2d(
|
||||
out_chans,
|
||||
out_chans,
|
||||
kernel_size=3,
|
||||
padding=1,
|
||||
bias=False,
|
||||
),
|
||||
LayerNorm2d(out_chans),
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = self.patch_embed(x)
|
||||
if self.pos_embed is not None:
|
||||
x = x + self.pos_embed
|
||||
|
||||
for blk in self.blocks:
|
||||
x = blk(x)
|
||||
|
||||
dtype = x.dtype
|
||||
if dtype == torch.float16: # prevent overflow
|
||||
with torch.autocast(device_type="cuda", dtype=torch.float32):
|
||||
x = self.neck(x.permute(0, 3, 1, 2))
|
||||
x = x.to(dtype)
|
||||
else:
|
||||
x = self.neck(x.permute(0, 3, 1, 2))
|
||||
return x
|
||||
|
||||
|
||||
class Block(nn.Module):
|
||||
"""Transformer blocks with support of window attention and residual propagation blocks"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_heads: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
qkv_bias: bool = True,
|
||||
norm_layer: Type[nn.Module] = nn.LayerNorm,
|
||||
act_layer: Type[nn.Module] = nn.GELU,
|
||||
use_rel_pos: bool = False,
|
||||
rel_pos_zero_init: bool = True,
|
||||
window_size: int = 0,
|
||||
input_size: Optional[Tuple[int, int]] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Args:
|
||||
dim (int): Number of input channels.
|
||||
num_heads (int): Number of attention heads in each ViT block.
|
||||
mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.
|
||||
qkv_bias (bool): If True, add a learnable bias to query, key, value.
|
||||
norm_layer (nn.Module): Normalization layer.
|
||||
act_layer (nn.Module): Activation layer.
|
||||
use_rel_pos (bool): If True, add relative positional embeddings to the attention map.
|
||||
rel_pos_zero_init (bool): If True, zero initialize relative positional parameters.
|
||||
window_size (int): Window size for window attention blocks. If it equals 0, then
|
||||
use global attention.
|
||||
input_size (tuple(int, int) or None): Input resolution for calculating the relative
|
||||
positional parameter size.
|
||||
"""
|
||||
super().__init__()
|
||||
self.norm1 = norm_layer(dim)
|
||||
self.attn = Attention(
|
||||
dim,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=qkv_bias,
|
||||
use_rel_pos=use_rel_pos,
|
||||
rel_pos_zero_init=rel_pos_zero_init,
|
||||
input_size=input_size if window_size == 0 else (window_size, window_size),
|
||||
)
|
||||
|
||||
self.norm2 = norm_layer(dim)
|
||||
self.mlp = MLPBlock(
|
||||
embedding_dim=dim, mlp_dim=int(dim * mlp_ratio), act=act_layer
|
||||
)
|
||||
|
||||
self.window_size = window_size
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
shortcut = x
|
||||
x = self.norm1(x)
|
||||
# Window partition
|
||||
if self.window_size > 0:
|
||||
H, W = x.shape[1], x.shape[2]
|
||||
x, pad_hw = window_partition(x, self.window_size)
|
||||
|
||||
x = self.attn(x)
|
||||
# Reverse window partition
|
||||
if self.window_size > 0:
|
||||
x = window_unpartition(x, self.window_size, pad_hw, (H, W))
|
||||
|
||||
x = shortcut + x
|
||||
x = x + self.mlp(self.norm2(x))
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
"""Multi-head Attention block with relative position embeddings."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_heads: int = 8,
|
||||
qkv_bias: bool = True,
|
||||
use_rel_pos: bool = False,
|
||||
rel_pos_zero_init: bool = True,
|
||||
input_size: Optional[Tuple[int, int]] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Args:
|
||||
dim (int): Number of input channels.
|
||||
num_heads (int): Number of attention heads.
|
||||
qkv_bias (bool): If True, add a learnable bias to query, key, value.
|
||||
rel_pos (bool): If True, add relative positional embeddings to the attention map.
|
||||
rel_pos_zero_init (bool): If True, zero initialize relative positional parameters.
|
||||
input_size (tuple(int, int) or None): Input resolution for calculating the relative
|
||||
positional parameter size.
|
||||
"""
|
||||
super().__init__()
|
||||
self.num_heads = num_heads
|
||||
head_dim = dim // num_heads
|
||||
self.scale = head_dim**-0.5
|
||||
|
||||
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
|
||||
self.proj = nn.Linear(dim, dim)
|
||||
|
||||
self.use_rel_pos = use_rel_pos
|
||||
if self.use_rel_pos:
|
||||
assert (
|
||||
input_size is not None
|
||||
), "Input size must be provided if using relative positional encoding."
|
||||
# initialize relative positional embeddings
|
||||
self.rel_pos_h = nn.Parameter(torch.zeros(2 * input_size[0] - 1, head_dim))
|
||||
self.rel_pos_w = nn.Parameter(torch.zeros(2 * input_size[1] - 1, head_dim))
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
B, H, W, _ = x.shape
|
||||
# qkv with shape (3, B, nHead, H * W, C)
|
||||
qkv = (
|
||||
self.qkv(x).reshape(B, H * W, 3, self.num_heads, -1).permute(2, 0, 3, 1, 4)
|
||||
)
|
||||
# q, k, v with shape (B * nHead, H * W, C)
|
||||
q, k, v = qkv.reshape(3, B * self.num_heads, H * W, -1).unbind(0)
|
||||
|
||||
attn = (q * self.scale) @ k.transpose(-2, -1)
|
||||
|
||||
if self.use_rel_pos:
|
||||
attn = add_decomposed_rel_pos(
|
||||
attn, q, self.rel_pos_h, self.rel_pos_w, (H, W), (H, W)
|
||||
)
|
||||
|
||||
attn = attn.softmax(dim=-1)
|
||||
x = (
|
||||
(attn @ v)
|
||||
.view(B, self.num_heads, H, W, -1)
|
||||
.permute(0, 2, 3, 1, 4)
|
||||
.reshape(B, H, W, -1)
|
||||
)
|
||||
x = self.proj(x)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
def window_partition(
|
||||
x: torch.Tensor, window_size: int
|
||||
) -> Tuple[torch.Tensor, Tuple[int, int]]:
|
||||
"""
|
||||
Partition into non-overlapping windows with padding if needed.
|
||||
Args:
|
||||
x (tensor): input tokens with [B, H, W, C].
|
||||
window_size (int): window size.
|
||||
|
||||
Returns:
|
||||
windows: windows after partition with [B * num_windows, window_size, window_size, C].
|
||||
(Hp, Wp): padded height and width before partition
|
||||
"""
|
||||
B, H, W, C = x.shape
|
||||
|
||||
pad_h = (window_size - H % window_size) % window_size
|
||||
pad_w = (window_size - W % window_size) % window_size
|
||||
if pad_h > 0 or pad_w > 0:
|
||||
x = F.pad(x, (0, 0, 0, pad_w, 0, pad_h))
|
||||
Hp, Wp = H + pad_h, W + pad_w
|
||||
|
||||
x = x.view(B, Hp // window_size, window_size, Wp // window_size, window_size, C)
|
||||
windows = (
|
||||
x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C)
|
||||
)
|
||||
return windows, (Hp, Wp)
|
||||
|
||||
|
||||
def window_unpartition(
|
||||
windows: torch.Tensor,
|
||||
window_size: int,
|
||||
pad_hw: Tuple[int, int],
|
||||
hw: Tuple[int, int],
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Window unpartition into original sequences and removing padding.
|
||||
Args:
|
||||
windows (tensor): input tokens with [B * num_windows, window_size, window_size, C].
|
||||
window_size (int): window size.
|
||||
pad_hw (Tuple): padded height and width (Hp, Wp).
|
||||
hw (Tuple): original height and width (H, W) before padding.
|
||||
|
||||
Returns:
|
||||
x: unpartitioned sequences with [B, H, W, C].
|
||||
"""
|
||||
Hp, Wp = pad_hw
|
||||
H, W = hw
|
||||
B = windows.shape[0] // (Hp * Wp // window_size // window_size)
|
||||
x = windows.view(
|
||||
B, Hp // window_size, Wp // window_size, window_size, window_size, -1
|
||||
)
|
||||
x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, Hp, Wp, -1)
|
||||
|
||||
if Hp > H or Wp > W:
|
||||
x = x[:, :H, :W, :].contiguous()
|
||||
return x
|
||||
|
||||
|
||||
def get_rel_pos(q_size: int, k_size: int, rel_pos: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Get relative positional embeddings according to the relative positions of
|
||||
query and key sizes.
|
||||
Args:
|
||||
q_size (int): size of query q.
|
||||
k_size (int): size of key k.
|
||||
rel_pos (Tensor): relative position embeddings (L, C).
|
||||
|
||||
Returns:
|
||||
Extracted positional embeddings according to relative positions.
|
||||
"""
|
||||
max_rel_dist = int(2 * max(q_size, k_size) - 1)
|
||||
# Interpolate rel pos if needed.
|
||||
if rel_pos.shape[0] != max_rel_dist:
|
||||
# Interpolate rel pos.
|
||||
rel_pos_resized = F.interpolate(
|
||||
rel_pos.reshape(1, rel_pos.shape[0], -1).permute(0, 2, 1),
|
||||
size=max_rel_dist,
|
||||
mode="linear",
|
||||
)
|
||||
rel_pos_resized = rel_pos_resized.reshape(-1, max_rel_dist).permute(1, 0)
|
||||
else:
|
||||
rel_pos_resized = rel_pos
|
||||
|
||||
# Scale the coords with short length if shapes for q and k are different.
|
||||
q_coords = torch.arange(q_size)[:, None] * max(k_size / q_size, 1.0)
|
||||
k_coords = torch.arange(k_size)[None, :] * max(q_size / k_size, 1.0)
|
||||
relative_coords = (q_coords - k_coords) + (k_size - 1) * max(q_size / k_size, 1.0)
|
||||
|
||||
return rel_pos_resized[relative_coords.long()]
|
||||
|
||||
|
||||
def add_decomposed_rel_pos(
|
||||
attn: torch.Tensor,
|
||||
q: torch.Tensor,
|
||||
rel_pos_h: torch.Tensor,
|
||||
rel_pos_w: torch.Tensor,
|
||||
q_size: Tuple[int, int],
|
||||
k_size: Tuple[int, int],
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Calculate decomposed Relative Positional Embeddings from :paper:`mvitv2`.
|
||||
https://github.com/facebookresearch/mvit/blob/19786631e330df9f3622e5402b4a419a263a2c80/mvit/models/attention.py # noqa B950
|
||||
Args:
|
||||
attn (Tensor): attention map.
|
||||
q (Tensor): query q in the attention layer with shape (B, q_h * q_w, C).
|
||||
rel_pos_h (Tensor): relative position embeddings (Lh, C) for height axis.
|
||||
rel_pos_w (Tensor): relative position embeddings (Lw, C) for width axis.
|
||||
q_size (Tuple): spatial sequence size of query q with (q_h, q_w).
|
||||
k_size (Tuple): spatial sequence size of key k with (k_h, k_w).
|
||||
|
||||
Returns:
|
||||
attn (Tensor): attention map with added relative positional embeddings.
|
||||
"""
|
||||
q_h, q_w = q_size
|
||||
k_h, k_w = k_size
|
||||
Rh = get_rel_pos(q_h, k_h, rel_pos_h)
|
||||
Rw = get_rel_pos(q_w, k_w, rel_pos_w)
|
||||
|
||||
B, _, dim = q.shape
|
||||
r_q = q.reshape(B, q_h, q_w, dim)
|
||||
rel_h = torch.einsum("bhwc,hkc->bhwk", r_q, Rh)
|
||||
rel_w = torch.einsum("bhwc,wkc->bhwk", r_q, Rw)
|
||||
|
||||
attn = (
|
||||
attn.view(B, q_h, q_w, k_h, k_w)
|
||||
+ rel_h[:, :, :, :, None]
|
||||
+ rel_w[:, :, :, None, :]
|
||||
).view(B, q_h * q_w, k_h * k_w)
|
||||
|
||||
return attn
|
||||
|
||||
|
||||
class PatchEmbed(nn.Module):
|
||||
"""
|
||||
Image to Patch Embedding.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
kernel_size: Tuple[int, int] = (16, 16),
|
||||
stride: Tuple[int, int] = (16, 16),
|
||||
padding: Tuple[int, int] = (0, 0),
|
||||
in_chans: int = 3,
|
||||
embed_dim: int = 768,
|
||||
) -> None:
|
||||
"""
|
||||
Args:
|
||||
kernel_size (Tuple): kernel size of the projection layer.
|
||||
stride (Tuple): stride of the projection layer.
|
||||
padding (Tuple): padding size of the projection layer.
|
||||
in_chans (int): Number of input image channels.
|
||||
embed_dim (int): Patch embedding dimension.
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
self.proj = nn.Conv2d(
|
||||
in_chans, embed_dim, kernel_size=kernel_size, stride=stride, padding=padding
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = self.proj(x)
|
||||
# B C H W -> B H W C
|
||||
x = x.permute(0, 2, 3, 1)
|
||||
return x
|
||||
@@ -0,0 +1,191 @@
|
||||
# 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.
|
||||
|
||||
from typing import List, Tuple, Type
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.nn import functional as F
|
||||
|
||||
from .common import LayerNorm2d
|
||||
|
||||
|
||||
class MaskDecoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
transformer_dim: int,
|
||||
transformer: nn.Module,
|
||||
num_multimask_outputs: int = 3,
|
||||
activation: Type[nn.Module] = nn.GELU,
|
||||
iou_head_depth: int = 3,
|
||||
iou_head_hidden_dim: int = 256,
|
||||
) -> None:
|
||||
"""
|
||||
Predicts masks given an image and prompt embeddings, using a
|
||||
transformer architecture.
|
||||
|
||||
Arguments:
|
||||
transformer_dim (int): the channel dimension of the transformer
|
||||
transformer (nn.Module): the transformer used to predict masks
|
||||
num_multimask_outputs (int): the number of masks to predict
|
||||
when disambiguating masks
|
||||
activation (nn.Module): the type of activation to use when
|
||||
upscaling masks
|
||||
iou_head_depth (int): the depth of the MLP used to predict
|
||||
mask quality
|
||||
iou_head_hidden_dim (int): the hidden dimension of the MLP
|
||||
used to predict mask quality
|
||||
"""
|
||||
super().__init__()
|
||||
self.transformer_dim = transformer_dim
|
||||
self.transformer = transformer
|
||||
|
||||
self.num_multimask_outputs = num_multimask_outputs
|
||||
|
||||
self.iou_token = nn.Embedding(1, transformer_dim)
|
||||
self.num_mask_tokens = num_multimask_outputs + 1
|
||||
self.mask_tokens = nn.Embedding(self.num_mask_tokens, transformer_dim)
|
||||
|
||||
self.output_upscaling = nn.Sequential(
|
||||
nn.ConvTranspose2d(
|
||||
transformer_dim, transformer_dim // 4, kernel_size=2, stride=2
|
||||
),
|
||||
LayerNorm2d(transformer_dim // 4),
|
||||
activation(),
|
||||
nn.ConvTranspose2d(
|
||||
transformer_dim // 4, transformer_dim // 8, kernel_size=2, stride=2
|
||||
),
|
||||
activation(),
|
||||
)
|
||||
self.output_hypernetworks_mlps = nn.ModuleList(
|
||||
[
|
||||
MLP(transformer_dim, transformer_dim, transformer_dim // 8, 3)
|
||||
for i in range(self.num_mask_tokens)
|
||||
]
|
||||
)
|
||||
|
||||
self.iou_prediction_head = MLP(
|
||||
transformer_dim, iou_head_hidden_dim, self.num_mask_tokens, iou_head_depth
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
image_embeddings: torch.Tensor,
|
||||
image_pe: torch.Tensor,
|
||||
sparse_prompt_embeddings: torch.Tensor,
|
||||
dense_prompt_embeddings: torch.Tensor,
|
||||
multimask_output: bool,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Predict masks given image and prompt embeddings.
|
||||
|
||||
Arguments:
|
||||
image_embeddings (torch.Tensor): the embeddings from the image encoder
|
||||
image_pe (torch.Tensor): positional encoding with the shape of image_embeddings
|
||||
sparse_prompt_embeddings (torch.Tensor): the embeddings of the points and boxes
|
||||
dense_prompt_embeddings (torch.Tensor): the embeddings of the mask inputs
|
||||
multimask_output (bool): Whether to return multiple masks or a single
|
||||
mask.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: batched predicted masks
|
||||
torch.Tensor: batched predictions of mask quality
|
||||
"""
|
||||
masks, iou_pred = self.predict_masks(
|
||||
image_embeddings=image_embeddings,
|
||||
image_pe=image_pe,
|
||||
sparse_prompt_embeddings=sparse_prompt_embeddings,
|
||||
dense_prompt_embeddings=dense_prompt_embeddings,
|
||||
)
|
||||
|
||||
# Select the correct mask or masks for output
|
||||
if multimask_output:
|
||||
mask_slice = slice(1, None)
|
||||
else:
|
||||
mask_slice = slice(0, 1)
|
||||
masks = masks[:, mask_slice, :, :]
|
||||
iou_pred = iou_pred[:, mask_slice]
|
||||
|
||||
# Prepare output
|
||||
return masks, iou_pred
|
||||
|
||||
def predict_masks(
|
||||
self,
|
||||
image_embeddings: torch.Tensor,
|
||||
image_pe: torch.Tensor,
|
||||
sparse_prompt_embeddings: torch.Tensor,
|
||||
dense_prompt_embeddings: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Predicts masks. See 'forward' for more details."""
|
||||
# Concatenate output tokens
|
||||
output_tokens = torch.cat(
|
||||
[self.iou_token.weight, self.mask_tokens.weight], dim=0
|
||||
)
|
||||
output_tokens = output_tokens.unsqueeze(0).expand(
|
||||
sparse_prompt_embeddings.size(0), -1, -1
|
||||
)
|
||||
|
||||
tokens = torch.cat((output_tokens, sparse_prompt_embeddings), dim=1)
|
||||
|
||||
# image_embeddings: [1, C, H, W], tokens: [B, N, C]
|
||||
# dense_prompt_embeddings: [B, C, H, W]
|
||||
# Expand per-image data in batch direction to be per-mask
|
||||
src = torch.repeat_interleave(image_embeddings, tokens.shape[0], dim=0)
|
||||
src = src + dense_prompt_embeddings
|
||||
pos_src = torch.repeat_interleave(image_pe, tokens.shape[0], dim=0)
|
||||
b, c, h, w = src.shape
|
||||
|
||||
# Run the transformer
|
||||
hs, src = self.transformer(src, pos_src, tokens)
|
||||
iou_token_out = hs[:, 0, :]
|
||||
mask_tokens_out = hs[:, 1 : (1 + self.num_mask_tokens), :]
|
||||
|
||||
# Upscale mask embeddings and predict masks using the mask tokens
|
||||
src = src.transpose(1, 2).view(b, c, h, w)
|
||||
upscaled_embedding = self.output_upscaling(src)
|
||||
hyper_in_list: List[torch.Tensor] = []
|
||||
for i in range(self.num_mask_tokens):
|
||||
hyper_in_list.append(
|
||||
self.output_hypernetworks_mlps[i](mask_tokens_out[:, i, :])
|
||||
)
|
||||
hyper_in = torch.stack(hyper_in_list, dim=1)
|
||||
b, c, h, w = upscaled_embedding.shape
|
||||
masks = (hyper_in @ upscaled_embedding.view(b, c, h * w)).view(
|
||||
b, self.num_mask_tokens, h, w
|
||||
)
|
||||
|
||||
# Generate mask quality predictions
|
||||
iou_pred = self.iou_prediction_head(iou_token_out)
|
||||
|
||||
return masks, iou_pred
|
||||
|
||||
|
||||
# Lightly adapted from
|
||||
# https://github.com/facebookresearch/MaskFormer/blob/main/mask_former/modeling/transformer/transformer_predictor.py # noqa
|
||||
class MLP(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_dim: int,
|
||||
hidden_dim: int,
|
||||
output_dim: int,
|
||||
num_layers: int,
|
||||
sigmoid_output: bool = False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.num_layers = num_layers
|
||||
h = [hidden_dim] * (num_layers - 1)
|
||||
self.layers = nn.ModuleList(
|
||||
nn.Linear(n, k) for n, k in zip([input_dim] + h, h + [output_dim])
|
||||
)
|
||||
self.sigmoid_output = sigmoid_output
|
||||
|
||||
def forward(self, x):
|
||||
for i, layer in enumerate(self.layers):
|
||||
x = F.relu(layer(x)) if i < self.num_layers - 1 else layer(x)
|
||||
if self.sigmoid_output:
|
||||
x = F.sigmoid(x)
|
||||
return x
|
||||
@@ -0,0 +1,238 @@
|
||||
# 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.
|
||||
|
||||
from typing import Any, Optional, Tuple, Type
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from .common import LayerNorm2d
|
||||
|
||||
|
||||
class PromptEncoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
embed_dim: int,
|
||||
image_embedding_size: Tuple[int, int],
|
||||
input_image_size: Tuple[int, int],
|
||||
mask_in_chans: int,
|
||||
activation: Type[nn.Module] = nn.GELU,
|
||||
) -> None:
|
||||
"""
|
||||
Encodes prompts for input to SAM's mask decoder.
|
||||
|
||||
Arguments:
|
||||
embed_dim (int): The prompts' embedding dimension
|
||||
image_embedding_size (tuple(int, int)): The spatial size of the
|
||||
image embedding, as (H, W).
|
||||
input_image_size (int): The padded size of the image as input
|
||||
to the image encoder, as (H, W).
|
||||
mask_in_chans (int): The number of hidden channels used for
|
||||
encoding input masks.
|
||||
activation (nn.Module): The activation to use when encoding
|
||||
input masks.
|
||||
"""
|
||||
super().__init__()
|
||||
self.embed_dim = embed_dim
|
||||
self.input_image_size = input_image_size
|
||||
self.image_embedding_size = image_embedding_size
|
||||
self.pe_layer = PositionEmbeddingRandom(embed_dim // 2)
|
||||
|
||||
self.num_point_embeddings: int = 4 # pos/neg point + 2 box corners
|
||||
point_embeddings = [
|
||||
nn.Embedding(1, embed_dim) for i in range(self.num_point_embeddings)
|
||||
]
|
||||
self.point_embeddings = nn.ModuleList(point_embeddings)
|
||||
self.not_a_point_embed = nn.Embedding(1, embed_dim)
|
||||
|
||||
self.mask_input_size = (
|
||||
4 * image_embedding_size[0],
|
||||
4 * image_embedding_size[1],
|
||||
)
|
||||
self.mask_downscaling = nn.Sequential(
|
||||
nn.Conv2d(1, mask_in_chans // 4, kernel_size=2, stride=2),
|
||||
LayerNorm2d(mask_in_chans // 4),
|
||||
activation(),
|
||||
nn.Conv2d(mask_in_chans // 4, mask_in_chans, kernel_size=2, stride=2),
|
||||
LayerNorm2d(mask_in_chans),
|
||||
activation(),
|
||||
nn.Conv2d(mask_in_chans, embed_dim, kernel_size=1),
|
||||
)
|
||||
self.no_mask_embed = nn.Embedding(1, embed_dim)
|
||||
|
||||
def get_dense_pe(self) -> torch.Tensor:
|
||||
"""
|
||||
Returns the positional encoding used to encode point prompts,
|
||||
applied to a dense set of points the shape of the image encoding.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Positional encoding with shape
|
||||
1x(embed_dim)x(embedding_h)x(embedding_w)
|
||||
"""
|
||||
return self.pe_layer(self.image_embedding_size).unsqueeze(0)
|
||||
|
||||
def _embed_points(
|
||||
self,
|
||||
points: torch.Tensor,
|
||||
labels: torch.Tensor,
|
||||
pad: bool,
|
||||
) -> torch.Tensor:
|
||||
"""Embeds point prompts."""
|
||||
points = points + 0.5 # Shift to center of pixel
|
||||
if pad:
|
||||
padding_point = torch.zeros((points.shape[0], 1, 2), device=points.device)
|
||||
padding_label = -torch.ones((labels.shape[0], 1), device=labels.device)
|
||||
points = torch.cat([points, padding_point], dim=1)
|
||||
labels = torch.cat([labels, padding_label], dim=1)
|
||||
point_embedding = self.pe_layer.forward_with_coords(
|
||||
points, self.input_image_size
|
||||
)
|
||||
point_embedding[labels == -1] = 0.0
|
||||
point_embedding[labels == -1] += self.not_a_point_embed.weight
|
||||
point_embedding[labels == 0] += self.point_embeddings[0].weight
|
||||
point_embedding[labels == 1] += self.point_embeddings[1].weight
|
||||
return point_embedding
|
||||
|
||||
def _embed_boxes(self, boxes: torch.Tensor) -> torch.Tensor:
|
||||
"""Embeds box prompts."""
|
||||
boxes = boxes + 0.5 # Shift to center of pixel
|
||||
coords = boxes.reshape(-1, 2, 2)
|
||||
corner_embedding = self.pe_layer.forward_with_coords(
|
||||
coords, self.input_image_size
|
||||
)
|
||||
corner_embedding[:, 0, :] += self.point_embeddings[2].weight
|
||||
corner_embedding[:, 1, :] += self.point_embeddings[3].weight
|
||||
return corner_embedding
|
||||
|
||||
def _embed_masks(self, masks: torch.Tensor) -> torch.Tensor:
|
||||
"""Embeds mask inputs."""
|
||||
mask_embedding = self.mask_downscaling(masks)
|
||||
return mask_embedding
|
||||
|
||||
def _get_batch_size(
|
||||
self,
|
||||
points: Optional[Tuple[torch.Tensor, torch.Tensor]],
|
||||
boxes: Optional[torch.Tensor],
|
||||
masks: Optional[torch.Tensor],
|
||||
text_embeds: Optional[torch.Tensor],
|
||||
) -> int:
|
||||
"""
|
||||
Gets the batch size of the output given the batch size of the input prompts.
|
||||
"""
|
||||
if points is not None:
|
||||
return points[0].shape[0]
|
||||
elif boxes is not None:
|
||||
return boxes.shape[0]
|
||||
elif masks is not None:
|
||||
return masks.shape[0]
|
||||
elif text_embeds is not None:
|
||||
return text_embeds.shape[0]
|
||||
else:
|
||||
return 1
|
||||
|
||||
def _get_device(self) -> torch.device:
|
||||
return self.point_embeddings[0].weight.device
|
||||
|
||||
def forward(
|
||||
self,
|
||||
points: Optional[Tuple[torch.Tensor, torch.Tensor]],
|
||||
boxes: Optional[torch.Tensor],
|
||||
masks: Optional[torch.Tensor],
|
||||
text_embeds: Optional[torch.Tensor],
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Embeds different types of prompts, returning both sparse and dense
|
||||
embeddings.
|
||||
|
||||
Arguments:
|
||||
points (tuple(torch.Tensor, torch.Tensor) or none): point coordinates
|
||||
and labels to embed.
|
||||
boxes (torch.Tensor or none): boxes to embed
|
||||
masks (torch.Tensor or none): masks to embed
|
||||
|
||||
Returns:
|
||||
torch.Tensor: sparse embeddings for the points and boxes, with shape
|
||||
BxNx(embed_dim), where N is determined by the number of input points
|
||||
and boxes.
|
||||
torch.Tensor: dense embeddings for the masks, in the shape
|
||||
Bx(embed_dim)x(embed_H)x(embed_W)
|
||||
"""
|
||||
bs = self._get_batch_size(points, boxes, masks, text_embeds)
|
||||
sparse_embeddings = torch.empty(
|
||||
(bs, 0, self.embed_dim), device=self._get_device()
|
||||
)
|
||||
if points is not None:
|
||||
coords, labels = points
|
||||
point_embeddings = self._embed_points(coords, labels, pad=(boxes is None))
|
||||
sparse_embeddings = torch.cat([sparse_embeddings, point_embeddings], dim=1)
|
||||
if boxes is not None:
|
||||
box_embeddings = self._embed_boxes(boxes)
|
||||
sparse_embeddings = torch.cat([sparse_embeddings, box_embeddings], dim=1)
|
||||
|
||||
if text_embeds is not None:
|
||||
sparse_embeddings = torch.cat([sparse_embeddings, text_embeds], dim=1)
|
||||
|
||||
if masks is not None:
|
||||
dense_embeddings = self._embed_masks(masks)
|
||||
else:
|
||||
dense_embeddings = self.no_mask_embed.weight.reshape(1, -1, 1, 1).expand(
|
||||
bs, -1, self.image_embedding_size[0], self.image_embedding_size[1]
|
||||
)
|
||||
|
||||
return sparse_embeddings, dense_embeddings
|
||||
|
||||
|
||||
class PositionEmbeddingRandom(nn.Module):
|
||||
"""
|
||||
Positional encoding using random spatial frequencies.
|
||||
"""
|
||||
|
||||
def __init__(self, num_pos_feats: int = 64, scale: Optional[float] = None) -> None:
|
||||
super().__init__()
|
||||
if scale is None or scale <= 0.0:
|
||||
scale = 1.0
|
||||
self.register_buffer(
|
||||
"positional_encoding_gaussian_matrix",
|
||||
scale * torch.randn((2, num_pos_feats)),
|
||||
)
|
||||
|
||||
def _pe_encoding(self, coords: torch.Tensor) -> torch.Tensor:
|
||||
"""Positionally encode points that are normalized to [0,1]."""
|
||||
# assuming coords are in [0, 1]^2 square and have d_1 x ... x d_n x 2 shape
|
||||
coords = 2 * coords - 1
|
||||
|
||||
if coords.dtype != self.positional_encoding_gaussian_matrix.dtype:
|
||||
coords = coords.to(self.positional_encoding_gaussian_matrix.dtype)
|
||||
|
||||
coords = coords @ self.positional_encoding_gaussian_matrix
|
||||
coords = 2 * np.pi * coords
|
||||
# outputs d_1 x ... x d_n x C shape
|
||||
return torch.cat([torch.sin(coords), torch.cos(coords)], dim=-1)
|
||||
|
||||
def forward(self, size: Tuple[int, int]) -> torch.Tensor:
|
||||
"""Generate positional encoding for a grid of the specified size."""
|
||||
h, w = size
|
||||
device: Any = self.positional_encoding_gaussian_matrix.device
|
||||
grid = torch.ones(
|
||||
(h, w), device=device, dtype=self.positional_encoding_gaussian_matrix.dtype
|
||||
)
|
||||
y_embed = grid.cumsum(dim=0) - 0.5
|
||||
x_embed = grid.cumsum(dim=1) - 0.5
|
||||
y_embed = y_embed / h
|
||||
x_embed = x_embed / w
|
||||
|
||||
pe = self._pe_encoding(torch.stack([x_embed, y_embed], dim=-1))
|
||||
return pe.permute(2, 0, 1) # C x H x W
|
||||
|
||||
def forward_with_coords(
|
||||
self, coords_input: torch.Tensor, image_size: Tuple[int, int]
|
||||
) -> torch.Tensor:
|
||||
"""Positionally encode points that are not normalized to [0,1]."""
|
||||
coords = coords_input.clone()
|
||||
coords[:, :, 0] = coords[:, :, 0] / image_size[1]
|
||||
coords[:, :, 1] = coords[:, :, 1] / image_size[0]
|
||||
return self._pe_encoding(coords.to(torch.float)) # B x N x C
|
||||
@@ -0,0 +1,184 @@
|
||||
# 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.
|
||||
|
||||
from typing import Any, Dict, List, Tuple
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.nn import functional as F
|
||||
|
||||
from .image_encoder import ImageEncoderViT
|
||||
from .mask_decoder import MaskDecoder
|
||||
from .prompt_encoder import PromptEncoder
|
||||
|
||||
|
||||
class Sam(nn.Module):
|
||||
mask_threshold: float = 0.0
|
||||
image_format: str = "RGB"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
image_encoder: ImageEncoderViT,
|
||||
prompt_encoder: PromptEncoder,
|
||||
mask_decoder: MaskDecoder,
|
||||
pixel_mean: List[float] = [123.675, 116.28, 103.53],
|
||||
pixel_std: List[float] = [58.395, 57.12, 57.375],
|
||||
) -> None:
|
||||
"""
|
||||
SAM predicts object masks from an image and input prompts.
|
||||
|
||||
Arguments:
|
||||
image_encoder (ImageEncoderViT): The backbone used to encode the
|
||||
image into image embeddings that allow for efficient mask prediction.
|
||||
prompt_encoder (PromptEncoder): Encodes various types of input prompts.
|
||||
mask_decoder (MaskDecoder): Predicts masks from the image embeddings
|
||||
and encoded prompts.
|
||||
pixel_mean (list(float)): Mean values for normalizing pixels in the input image.
|
||||
pixel_std (list(float)): Std values for normalizing pixels in the input image.
|
||||
"""
|
||||
super().__init__()
|
||||
self.image_encoder = image_encoder
|
||||
self.prompt_encoder = prompt_encoder
|
||||
self.mask_decoder = mask_decoder
|
||||
self.register_buffer(
|
||||
"pixel_mean", torch.Tensor(pixel_mean).view(-1, 1, 1), False
|
||||
)
|
||||
self.register_buffer("pixel_std", torch.Tensor(pixel_std).view(-1, 1, 1), False)
|
||||
|
||||
@property
|
||||
def device(self) -> Any:
|
||||
return self.pixel_mean.device
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
batched_input: List[Dict[str, Any]],
|
||||
multimask_output: bool,
|
||||
) -> List[Dict[str, torch.Tensor]]:
|
||||
"""
|
||||
Predicts masks end-to-end from provided images and prompts.
|
||||
If prompts are not known in advance, using SamPredictor is
|
||||
recommended over calling the model directly.
|
||||
|
||||
Arguments:
|
||||
batched_input (list(dict)): A list over input images, each a
|
||||
dictionary with the following keys. A prompt key can be
|
||||
excluded if it is not present.
|
||||
'image': The image as a torch tensor in 3xHxW format,
|
||||
already transformed for input to the model.
|
||||
'original_size': (tuple(int, int)) The original size of
|
||||
the image before transformation, as (H, W).
|
||||
'point_coords': (torch.Tensor) Batched point prompts for
|
||||
this image, with shape BxNx2. Already transformed to the
|
||||
input frame of the model.
|
||||
'point_labels': (torch.Tensor) Batched labels for point prompts,
|
||||
with shape BxN.
|
||||
'boxes': (torch.Tensor) Batched box inputs, with shape Bx4.
|
||||
Already transformed to the input frame of the model.
|
||||
'mask_inputs': (torch.Tensor) Batched mask inputs to the model,
|
||||
in the form Bx1xHxW.
|
||||
multimask_output (bool): Whether the model should predict multiple
|
||||
disambiguating masks, or return a single mask.
|
||||
|
||||
Returns:
|
||||
(list(dict)): A list over input images, where each element is
|
||||
as dictionary with the following keys.
|
||||
'masks': (torch.Tensor) Batched binary mask predictions,
|
||||
with shape BxCxHxW, where B is the number of input prompts,
|
||||
C is determined by multimask_output, and (H, W) is the
|
||||
original size of the image.
|
||||
'iou_predictions': (torch.Tensor) The model's predictions
|
||||
of mask quality, in shape BxC.
|
||||
'low_res_logits': (torch.Tensor) Low resolution logits with
|
||||
shape BxCxHxW, where H=W=256. Can be passed as mask input
|
||||
to subsequent iterations of prediction.
|
||||
"""
|
||||
input_images = torch.stack(
|
||||
[self.preprocess(x["image"]) for x in batched_input], dim=0
|
||||
)
|
||||
image_embeddings = self.image_encoder(input_images)
|
||||
|
||||
outputs = []
|
||||
for image_record, curr_embedding in zip(batched_input, image_embeddings):
|
||||
if "point_coords" in image_record:
|
||||
points = (image_record["point_coords"], image_record["point_labels"])
|
||||
else:
|
||||
points = None
|
||||
sparse_embeddings, dense_embeddings = self.prompt_encoder(
|
||||
points=points,
|
||||
boxes=image_record.get("boxes", None),
|
||||
masks=image_record.get("mask_inputs", None),
|
||||
)
|
||||
low_res_masks, iou_predictions = self.mask_decoder(
|
||||
image_embeddings=curr_embedding.unsqueeze(0),
|
||||
image_pe=self.prompt_encoder.get_dense_pe(),
|
||||
sparse_prompt_embeddings=sparse_embeddings,
|
||||
dense_prompt_embeddings=dense_embeddings,
|
||||
multimask_output=multimask_output,
|
||||
)
|
||||
masks = self.postprocess_masks(
|
||||
low_res_masks,
|
||||
input_size=image_record["image"].shape[-2:],
|
||||
original_size=image_record["original_size"],
|
||||
)
|
||||
masks = masks > self.mask_threshold
|
||||
outputs.append(
|
||||
{
|
||||
"masks": masks,
|
||||
"iou_predictions": iou_predictions,
|
||||
"low_res_logits": low_res_masks,
|
||||
}
|
||||
)
|
||||
return outputs
|
||||
|
||||
def postprocess_masks(
|
||||
self,
|
||||
masks: torch.Tensor,
|
||||
input_size: Tuple[int, ...],
|
||||
original_size: Tuple[int, ...],
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Remove padding and upscale masks to the original image size.
|
||||
|
||||
Arguments:
|
||||
masks (torch.Tensor): Batched masks from the mask_decoder,
|
||||
in BxCxHxW format.
|
||||
input_size (tuple(int, int)): The size of the image input to the
|
||||
model, in (H, W) format. Used to remove padding.
|
||||
original_size (tuple(int, int)): The original size of the image
|
||||
before resizing for input to the model, in (H, W) format.
|
||||
|
||||
Returns:
|
||||
(torch.Tensor): Batched masks in BxCxHxW format, where (H, W)
|
||||
is given by original_size.
|
||||
"""
|
||||
|
||||
dtype = masks.dtype
|
||||
|
||||
masks = F.interpolate(
|
||||
masks.float(),
|
||||
(self.image_encoder.img_size, self.image_encoder.img_size),
|
||||
mode="bilinear",
|
||||
align_corners=False,
|
||||
)
|
||||
# masks = masks.to(dtype)
|
||||
masks = masks[..., : input_size[0], : input_size[1]]
|
||||
masks = F.interpolate(
|
||||
masks, original_size, mode="bilinear", align_corners=False
|
||||
)
|
||||
return masks
|
||||
|
||||
def preprocess(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Normalize pixel values and pad to a square input."""
|
||||
# Normalize colors
|
||||
x = (x - self.pixel_mean) / self.pixel_std
|
||||
|
||||
# Pad
|
||||
h, w = x.shape[-2:]
|
||||
padh = self.image_encoder.img_size - h
|
||||
padw = self.image_encoder.img_size - w
|
||||
x = F.pad(x, (0, padw, 0, padh))
|
||||
return x
|
||||
@@ -0,0 +1,242 @@
|
||||
# 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 math
|
||||
from typing import Tuple, Type
|
||||
|
||||
import torch
|
||||
from torch import Tensor, nn
|
||||
|
||||
from .common import MLPBlock
|
||||
|
||||
|
||||
class TwoWayTransformer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
depth: int,
|
||||
embedding_dim: int,
|
||||
num_heads: int,
|
||||
mlp_dim: int,
|
||||
activation: Type[nn.Module] = nn.ReLU,
|
||||
attention_downsample_rate: int = 2,
|
||||
) -> None:
|
||||
"""
|
||||
A transformer decoder that attends to an input image using
|
||||
queries whose positional embedding is supplied.
|
||||
|
||||
Args:
|
||||
depth (int): number of layers in the transformer
|
||||
embedding_dim (int): the channel dimension for the input embeddings
|
||||
num_heads (int): the number of heads for multihead attention. Must
|
||||
divide embedding_dim
|
||||
mlp_dim (int): the channel dimension internal to the MLP block
|
||||
activation (nn.Module): the activation to use in the MLP block
|
||||
"""
|
||||
super().__init__()
|
||||
self.depth = depth
|
||||
self.embedding_dim = embedding_dim
|
||||
self.num_heads = num_heads
|
||||
self.mlp_dim = mlp_dim
|
||||
self.layers = nn.ModuleList()
|
||||
|
||||
for i in range(depth):
|
||||
self.layers.append(
|
||||
TwoWayAttentionBlock(
|
||||
embedding_dim=embedding_dim,
|
||||
num_heads=num_heads,
|
||||
mlp_dim=mlp_dim,
|
||||
activation=activation,
|
||||
attention_downsample_rate=attention_downsample_rate,
|
||||
skip_first_layer_pe=(i == 0),
|
||||
)
|
||||
)
|
||||
|
||||
self.final_attn_token_to_image = Attention(
|
||||
embedding_dim, num_heads, downsample_rate=attention_downsample_rate
|
||||
)
|
||||
self.norm_final_attn = nn.LayerNorm(embedding_dim)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
image_embedding: Tensor,
|
||||
image_pe: Tensor,
|
||||
point_embedding: Tensor,
|
||||
) -> Tuple[Tensor, Tensor]:
|
||||
"""
|
||||
Args:
|
||||
image_embedding (torch.Tensor): image to attend to. Should be shape
|
||||
B x embedding_dim x h x w for any h and w.
|
||||
image_pe (torch.Tensor): the positional encoding to add to the image. Must
|
||||
have the same shape as image_embedding.
|
||||
point_embedding (torch.Tensor): the embedding to add to the query points.
|
||||
Must have shape B x N_points x embedding_dim for any N_points.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: the processed point_embedding
|
||||
torch.Tensor: the processed image_embedding
|
||||
"""
|
||||
# BxCxHxW -> BxHWxC == B x N_image_tokens x C
|
||||
bs, c, h, w = image_embedding.shape
|
||||
image_embedding = image_embedding.flatten(2).permute(0, 2, 1)
|
||||
image_pe = image_pe.flatten(2).permute(0, 2, 1)
|
||||
|
||||
# Prepare queries
|
||||
queries = point_embedding
|
||||
keys = image_embedding
|
||||
|
||||
# Apply transformer blocks and final layernorm
|
||||
for layer in self.layers:
|
||||
queries, keys = layer(
|
||||
queries=queries,
|
||||
keys=keys,
|
||||
query_pe=point_embedding,
|
||||
key_pe=image_pe,
|
||||
)
|
||||
|
||||
# Apply the final attention layer from the points to the image
|
||||
q = queries + point_embedding
|
||||
k = keys + image_pe
|
||||
attn_out = self.final_attn_token_to_image(q=q, k=k, v=keys)
|
||||
queries = queries + attn_out
|
||||
queries = self.norm_final_attn(queries)
|
||||
|
||||
return queries, keys
|
||||
|
||||
|
||||
class TwoWayAttentionBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
embedding_dim: int,
|
||||
num_heads: int,
|
||||
mlp_dim: int = 2048,
|
||||
activation: Type[nn.Module] = nn.ReLU,
|
||||
attention_downsample_rate: int = 2,
|
||||
skip_first_layer_pe: bool = False,
|
||||
) -> None:
|
||||
"""
|
||||
A transformer block with four layers: (1) self-attention of sparse
|
||||
inputs, (2) cross attention of sparse inputs to dense inputs, (3) mlp
|
||||
block on sparse inputs, and (4) cross attention of dense inputs to sparse
|
||||
inputs.
|
||||
|
||||
Arguments:
|
||||
embedding_dim (int): the channel dimension of the embeddings
|
||||
num_heads (int): the number of heads in the attention layers
|
||||
mlp_dim (int): the hidden dimension of the mlp block
|
||||
activation (nn.Module): the activation of the mlp block
|
||||
skip_first_layer_pe (bool): skip the PE on the first layer
|
||||
"""
|
||||
super().__init__()
|
||||
self.self_attn = Attention(embedding_dim, num_heads)
|
||||
self.norm1 = nn.LayerNorm(embedding_dim)
|
||||
|
||||
self.cross_attn_token_to_image = Attention(
|
||||
embedding_dim, num_heads, downsample_rate=attention_downsample_rate
|
||||
)
|
||||
self.norm2 = nn.LayerNorm(embedding_dim)
|
||||
|
||||
self.mlp = MLPBlock(embedding_dim, mlp_dim, activation)
|
||||
self.norm3 = nn.LayerNorm(embedding_dim)
|
||||
|
||||
self.norm4 = nn.LayerNorm(embedding_dim)
|
||||
self.cross_attn_image_to_token = Attention(
|
||||
embedding_dim, num_heads, downsample_rate=attention_downsample_rate
|
||||
)
|
||||
|
||||
self.skip_first_layer_pe = skip_first_layer_pe
|
||||
|
||||
def forward(
|
||||
self, queries: Tensor, keys: Tensor, query_pe: Tensor, key_pe: Tensor
|
||||
) -> Tuple[Tensor, Tensor]:
|
||||
# Self attention block
|
||||
if self.skip_first_layer_pe:
|
||||
queries = self.self_attn(q=queries, k=queries, v=queries)
|
||||
else:
|
||||
q = queries + query_pe
|
||||
attn_out = self.self_attn(q=q, k=q, v=queries)
|
||||
queries = queries + attn_out
|
||||
queries = self.norm1(queries)
|
||||
|
||||
# Cross attention block, tokens attending to image embedding
|
||||
q = queries + query_pe
|
||||
k = keys + key_pe
|
||||
attn_out = self.cross_attn_token_to_image(q=q, k=k, v=keys)
|
||||
queries = queries + attn_out
|
||||
queries = self.norm2(queries)
|
||||
|
||||
# MLP block
|
||||
mlp_out = self.mlp(queries)
|
||||
queries = queries + mlp_out
|
||||
queries = self.norm3(queries)
|
||||
|
||||
# Cross attention block, image embedding attending to tokens
|
||||
q = queries + query_pe
|
||||
k = keys + key_pe
|
||||
attn_out = self.cross_attn_image_to_token(q=k, k=q, v=queries)
|
||||
keys = keys + attn_out
|
||||
keys = self.norm4(keys)
|
||||
|
||||
return queries, keys
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
"""
|
||||
An attention layer that allows for downscaling the size of the embedding
|
||||
after projection to queries, keys, and values.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
embedding_dim: int,
|
||||
num_heads: int,
|
||||
downsample_rate: int = 1,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.embedding_dim = embedding_dim
|
||||
self.internal_dim = embedding_dim // downsample_rate
|
||||
self.num_heads = num_heads
|
||||
assert (
|
||||
self.internal_dim % num_heads == 0
|
||||
), "num_heads must divide embedding_dim."
|
||||
|
||||
self.q_proj = nn.Linear(embedding_dim, self.internal_dim)
|
||||
self.k_proj = nn.Linear(embedding_dim, self.internal_dim)
|
||||
self.v_proj = nn.Linear(embedding_dim, self.internal_dim)
|
||||
self.out_proj = nn.Linear(self.internal_dim, embedding_dim)
|
||||
|
||||
def _separate_heads(self, x: Tensor, num_heads: int) -> Tensor:
|
||||
b, n, c = x.shape
|
||||
x = x.reshape(b, n, num_heads, c // num_heads)
|
||||
return x.transpose(1, 2) # B x N_heads x N_tokens x C_per_head
|
||||
|
||||
def _recombine_heads(self, x: Tensor) -> Tensor:
|
||||
b, n_heads, n_tokens, c_per_head = x.shape
|
||||
x = x.transpose(1, 2)
|
||||
return x.reshape(b, n_tokens, n_heads * c_per_head) # B x N_tokens x C
|
||||
|
||||
def forward(self, q: Tensor, k: Tensor, v: Tensor) -> Tensor:
|
||||
# Input projections
|
||||
q = self.q_proj(q)
|
||||
k = self.k_proj(k)
|
||||
v = self.v_proj(v)
|
||||
|
||||
# Separate into heads
|
||||
q = self._separate_heads(q, self.num_heads)
|
||||
k = self._separate_heads(k, self.num_heads)
|
||||
v = self._separate_heads(v, self.num_heads)
|
||||
|
||||
# Attention
|
||||
_, _, _, c_per_head = q.shape
|
||||
attn = q @ k.permute(0, 1, 3, 2) # B x N_heads x N_tokens x N_tokens
|
||||
attn = attn / math.sqrt(c_per_head)
|
||||
attn = torch.softmax(attn, dim=-1)
|
||||
|
||||
# Get output
|
||||
out = attn @ v
|
||||
out = self._recombine_heads(out)
|
||||
out = self.out_proj(out)
|
||||
|
||||
return out
|
||||
@@ -0,0 +1,284 @@
|
||||
# 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.
|
||||
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from .modeling import Sam
|
||||
from .utils.transforms import ResizeLongestSide
|
||||
|
||||
|
||||
class SamPredictor:
|
||||
def __init__(
|
||||
self,
|
||||
sam_model: Sam,
|
||||
) -> None:
|
||||
"""
|
||||
Uses SAM to calculate the image embedding for an image, and then
|
||||
allow repeated, efficient mask prediction given prompts.
|
||||
|
||||
Arguments:
|
||||
sam_model (Sam): The model to use for mask prediction.
|
||||
"""
|
||||
super().__init__()
|
||||
self.model = sam_model
|
||||
self.transform = ResizeLongestSide(sam_model.image_encoder.img_size)
|
||||
self.reset_image()
|
||||
|
||||
def set_image(
|
||||
self,
|
||||
image: np.ndarray,
|
||||
image_format: str = "RGB",
|
||||
) -> None:
|
||||
"""
|
||||
Calculates the image embeddings for the provided image, allowing
|
||||
masks to be predicted with the 'predict' method.
|
||||
|
||||
Arguments:
|
||||
image (np.ndarray): The image for calculating masks. Expects an
|
||||
image in HWC uint8 format, with pixel values in [0, 255].
|
||||
image_format (str): The color format of the image, in ['RGB', 'BGR'].
|
||||
"""
|
||||
assert image_format in [
|
||||
"RGB",
|
||||
"BGR",
|
||||
], f"image_format must be in ['RGB', 'BGR'], is {image_format}."
|
||||
if image_format != self.model.image_format:
|
||||
image = image[..., ::-1]
|
||||
|
||||
# Transform the image to the form expected by the model
|
||||
input_image = self.transform.apply_image(image)
|
||||
input_image_torch = torch.as_tensor(input_image, device=self.device)
|
||||
input_image_torch = input_image_torch.permute(2, 0, 1).contiguous()[
|
||||
None, :, :, :
|
||||
]
|
||||
|
||||
self.set_torch_image(input_image_torch, image.shape[:2])
|
||||
|
||||
@torch.no_grad()
|
||||
def set_torch_image(
|
||||
self,
|
||||
transformed_image: torch.Tensor,
|
||||
original_image_size: Tuple[int, ...],
|
||||
) -> None:
|
||||
"""
|
||||
Calculates the image embeddings for the provided image, allowing
|
||||
masks to be predicted with the 'predict' method. Expects the input
|
||||
image to be already transformed to the format expected by the model.
|
||||
|
||||
Arguments:
|
||||
transformed_image (torch.Tensor): The input image, with shape
|
||||
1x3xHxW, which has been transformed with ResizeLongestSide.
|
||||
original_image_size (tuple(int, int)): The size of the image
|
||||
before transformation, in (H, W) format.
|
||||
"""
|
||||
assert (
|
||||
len(transformed_image.shape) == 4
|
||||
and transformed_image.shape[1] == 3
|
||||
and max(*transformed_image.shape[2:]) == self.model.image_encoder.img_size
|
||||
), f"set_torch_image input must be BCHW with long side {self.model.image_encoder.img_size}."
|
||||
self.reset_image()
|
||||
|
||||
self.original_size = original_image_size
|
||||
self.input_size = tuple(transformed_image.shape[-2:])
|
||||
input_image = self.model.preprocess(transformed_image)
|
||||
self.features = self.model.image_encoder(input_image)
|
||||
self.is_image_set = True
|
||||
|
||||
def predict(
|
||||
self,
|
||||
point_coords: Optional[np.ndarray] = None,
|
||||
point_labels: Optional[np.ndarray] = None,
|
||||
box: Optional[np.ndarray] = None,
|
||||
mask_input: Optional[np.ndarray] = None,
|
||||
multimask_output: bool = True,
|
||||
return_logits: bool = False,
|
||||
) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
|
||||
"""
|
||||
Predict masks for the given input prompts, using the currently set image.
|
||||
|
||||
Arguments:
|
||||
point_coords (np.ndarray or None): A Nx2 array of point prompts to the
|
||||
model. Each point is in (X,Y) in pixels.
|
||||
point_labels (np.ndarray or None): A length N array of labels for the
|
||||
point prompts. 1 indicates a foreground point and 0 indicates a
|
||||
background point.
|
||||
box (np.ndarray or None): A length 4 array given a box prompt to the
|
||||
model, in XYXY format.
|
||||
mask_input (np.ndarray): A low resolution mask input to the model, typically
|
||||
coming from a previous prediction iteration. Has form 1xHxW, where
|
||||
for SAM, H=W=256.
|
||||
multimask_output (bool): If true, the model will return three masks.
|
||||
For ambiguous input prompts (such as a single click), this will often
|
||||
produce better masks than a single prediction. If only a single
|
||||
mask is needed, the model's predicted quality score can be used
|
||||
to select the best mask. For non-ambiguous prompts, such as multiple
|
||||
input prompts, multimask_output=False can give better results.
|
||||
return_logits (bool): If true, returns un-thresholded masks logits
|
||||
instead of a binary mask.
|
||||
|
||||
Returns:
|
||||
(np.ndarray): The output masks in CxHxW format, where C is the
|
||||
number of masks, and (H, W) is the original image size.
|
||||
(np.ndarray): An array of length C containing the model's
|
||||
predictions for the quality of each mask.
|
||||
(np.ndarray): An array of shape CxHxW, where C is the number
|
||||
of masks and H=W=256. These low resolution logits can be passed to
|
||||
a subsequent iteration as mask input.
|
||||
"""
|
||||
if not self.is_image_set:
|
||||
raise RuntimeError(
|
||||
"An image must be set with .set_image(...) before mask prediction."
|
||||
)
|
||||
|
||||
# Transform input prompts
|
||||
coords_torch, labels_torch, box_torch, mask_input_torch = None, None, None, None
|
||||
if point_coords is not None:
|
||||
assert (
|
||||
point_labels is not None
|
||||
), "point_labels must be supplied if point_coords is supplied."
|
||||
point_coords = self.transform.apply_coords(point_coords, self.original_size)
|
||||
coords_torch = torch.as_tensor(
|
||||
point_coords, dtype=torch.float, device=self.device
|
||||
)
|
||||
labels_torch = torch.as_tensor(
|
||||
point_labels, dtype=torch.int, device=self.device
|
||||
)
|
||||
coords_torch, labels_torch = coords_torch[None, :, :], labels_torch[None, :]
|
||||
if box is not None:
|
||||
box = self.transform.apply_boxes(box, self.original_size)
|
||||
box_torch = torch.as_tensor(box, dtype=torch.float, device=self.device)
|
||||
box_torch = box_torch[None, :]
|
||||
if mask_input is not None:
|
||||
mask_input_torch = torch.as_tensor(
|
||||
mask_input, dtype=torch.float, device=self.device
|
||||
)
|
||||
mask_input_torch = mask_input_torch[None, :, :, :]
|
||||
|
||||
masks, iou_predictions, low_res_masks = self.predict_torch(
|
||||
coords_torch,
|
||||
labels_torch,
|
||||
box_torch,
|
||||
mask_input_torch,
|
||||
multimask_output,
|
||||
return_logits=return_logits,
|
||||
)
|
||||
|
||||
masks_np = masks[0].detach().cpu().numpy()
|
||||
iou_predictions_np = iou_predictions[0].detach().cpu().numpy()
|
||||
low_res_masks_np = low_res_masks[0].detach().cpu().numpy()
|
||||
return masks_np, iou_predictions_np, low_res_masks_np
|
||||
|
||||
@torch.no_grad()
|
||||
def predict_torch(
|
||||
self,
|
||||
point_coords: Optional[torch.Tensor],
|
||||
point_labels: Optional[torch.Tensor],
|
||||
boxes: Optional[torch.Tensor] = None,
|
||||
mask_input: Optional[torch.Tensor] = None,
|
||||
multimask_output: bool = True,
|
||||
return_logits: bool = False,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Predict masks for the given input prompts, using the currently set image.
|
||||
Input prompts are batched torch tensors and are expected to already be
|
||||
transformed to the input frame using ResizeLongestSide.
|
||||
|
||||
Arguments:
|
||||
point_coords (torch.Tensor or None): A BxNx2 array of point prompts to the
|
||||
model. Each point is in (X,Y) in pixels.
|
||||
point_labels (torch.Tensor or None): A BxN array of labels for the
|
||||
point prompts. 1 indicates a foreground point and 0 indicates a
|
||||
background point.
|
||||
boxes (np.ndarray or None): A Bx4 array given a box prompt to the
|
||||
model, in XYXY format.
|
||||
mask_input (np.ndarray): A low resolution mask input to the model, typically
|
||||
coming from a previous prediction iteration. Has form Bx1xHxW, where
|
||||
for SAM, H=W=256. Masks returned by a previous iteration of the
|
||||
predict method do not need further transformation.
|
||||
multimask_output (bool): If true, the model will return three masks.
|
||||
For ambiguous input prompts (such as a single click), this will often
|
||||
produce better masks than a single prediction. If only a single
|
||||
mask is needed, the model's predicted quality score can be used
|
||||
to select the best mask. For non-ambiguous prompts, such as multiple
|
||||
input prompts, multimask_output=False can give better results.
|
||||
return_logits (bool): If true, returns un-thresholded masks logits
|
||||
instead of a binary mask.
|
||||
|
||||
Returns:
|
||||
(torch.Tensor): The output masks in BxCxHxW format, where C is the
|
||||
number of masks, and (H, W) is the original image size.
|
||||
(torch.Tensor): An array of shape BxC containing the model's
|
||||
predictions for the quality of each mask.
|
||||
(torch.Tensor): An array of shape BxCxHxW, where C is the number
|
||||
of masks and H=W=256. These low res logits can be passed to
|
||||
a subsequent iteration as mask input.
|
||||
"""
|
||||
if not self.is_image_set:
|
||||
raise RuntimeError(
|
||||
"An image must be set with .set_image(...) before mask prediction."
|
||||
)
|
||||
|
||||
if point_coords is not None:
|
||||
points = (point_coords, point_labels)
|
||||
else:
|
||||
points = None
|
||||
|
||||
# Embed prompts
|
||||
sparse_embeddings, dense_embeddings = self.model.prompt_encoder(
|
||||
points=points,
|
||||
boxes=boxes,
|
||||
masks=mask_input,
|
||||
)
|
||||
|
||||
# Predict masks
|
||||
low_res_masks, iou_predictions = self.model.mask_decoder(
|
||||
image_embeddings=self.features,
|
||||
image_pe=self.model.prompt_encoder.get_dense_pe(),
|
||||
sparse_prompt_embeddings=sparse_embeddings,
|
||||
dense_prompt_embeddings=dense_embeddings,
|
||||
multimask_output=multimask_output,
|
||||
)
|
||||
|
||||
# Upscale the masks to the original image resolution
|
||||
masks = self.model.postprocess_masks(
|
||||
low_res_masks, self.input_size, self.original_size
|
||||
)
|
||||
|
||||
if not return_logits:
|
||||
masks = masks > self.model.mask_threshold
|
||||
|
||||
return masks, iou_predictions, low_res_masks
|
||||
|
||||
def get_image_embedding(self) -> torch.Tensor:
|
||||
"""
|
||||
Returns the image embeddings for the currently set image, with
|
||||
shape 1xCxHxW, where C is the embedding dimension and (H,W) are
|
||||
the embedding spatial dimension of SAM (typically C=256, H=W=64).
|
||||
"""
|
||||
if not self.is_image_set:
|
||||
raise RuntimeError(
|
||||
"An image must be set with .set_image(...) to generate an embedding."
|
||||
)
|
||||
assert (
|
||||
self.features is not None
|
||||
), "Features must exist if an image has been set."
|
||||
return self.features
|
||||
|
||||
@property
|
||||
def device(self) -> torch.device:
|
||||
return self.model.device
|
||||
|
||||
def reset_image(self) -> None:
|
||||
"""Resets the currently set image."""
|
||||
self.is_image_set = False
|
||||
self.features = None
|
||||
self.orig_h = None
|
||||
self.orig_w = None
|
||||
self.input_h = None
|
||||
self.input_w = None
|
||||
@@ -0,0 +1,5 @@
|
||||
# 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.
|
||||
@@ -0,0 +1,346 @@
|
||||
# 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 math
|
||||
from copy import deepcopy
|
||||
from itertools import product
|
||||
from typing import Any, Dict, Generator, ItemsView, List, Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
class MaskData:
|
||||
"""
|
||||
A structure for storing masks and their related data in batched format.
|
||||
Implements basic filtering and concatenation.
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs) -> None:
|
||||
for v in kwargs.values():
|
||||
assert isinstance(
|
||||
v, (list, np.ndarray, torch.Tensor)
|
||||
), "MaskData only supports list, numpy arrays, and torch tensors."
|
||||
self._stats = dict(**kwargs)
|
||||
|
||||
def __setitem__(self, key: str, item: Any) -> None:
|
||||
assert isinstance(
|
||||
item, (list, np.ndarray, torch.Tensor)
|
||||
), "MaskData only supports list, numpy arrays, and torch tensors."
|
||||
self._stats[key] = item
|
||||
|
||||
def __delitem__(self, key: str) -> None:
|
||||
del self._stats[key]
|
||||
|
||||
def __getitem__(self, key: str) -> Any:
|
||||
return self._stats[key]
|
||||
|
||||
def items(self) -> ItemsView[str, Any]:
|
||||
return self._stats.items()
|
||||
|
||||
def filter(self, keep: torch.Tensor) -> None:
|
||||
for k, v in self._stats.items():
|
||||
if v is None:
|
||||
self._stats[k] = None
|
||||
elif isinstance(v, torch.Tensor):
|
||||
self._stats[k] = v[torch.as_tensor(keep, device=v.device)]
|
||||
elif isinstance(v, np.ndarray):
|
||||
self._stats[k] = v[keep.detach().cpu().numpy()]
|
||||
elif isinstance(v, list) and keep.dtype == torch.bool:
|
||||
self._stats[k] = [a for i, a in enumerate(v) if keep[i]]
|
||||
elif isinstance(v, list):
|
||||
self._stats[k] = [v[i] for i in keep]
|
||||
else:
|
||||
raise TypeError(f"MaskData key {k} has an unsupported type {type(v)}.")
|
||||
|
||||
def cat(self, new_stats: "MaskData") -> None:
|
||||
for k, v in new_stats.items():
|
||||
if k not in self._stats or self._stats[k] is None:
|
||||
self._stats[k] = deepcopy(v)
|
||||
elif isinstance(v, torch.Tensor):
|
||||
self._stats[k] = torch.cat([self._stats[k], v], dim=0)
|
||||
elif isinstance(v, np.ndarray):
|
||||
self._stats[k] = np.concatenate([self._stats[k], v], axis=0)
|
||||
elif isinstance(v, list):
|
||||
self._stats[k] = self._stats[k] + deepcopy(v)
|
||||
else:
|
||||
raise TypeError(f"MaskData key {k} has an unsupported type {type(v)}.")
|
||||
|
||||
def to_numpy(self) -> None:
|
||||
for k, v in self._stats.items():
|
||||
if isinstance(v, torch.Tensor):
|
||||
self._stats[k] = v.detach().cpu().numpy()
|
||||
|
||||
|
||||
def is_box_near_crop_edge(
|
||||
boxes: torch.Tensor, crop_box: List[int], orig_box: List[int], atol: float = 20.0
|
||||
) -> torch.Tensor:
|
||||
"""Filter masks at the edge of a crop, but not at the edge of the original image."""
|
||||
crop_box_torch = torch.as_tensor(crop_box, dtype=torch.float, device=boxes.device)
|
||||
orig_box_torch = torch.as_tensor(orig_box, dtype=torch.float, device=boxes.device)
|
||||
boxes = uncrop_boxes_xyxy(boxes, crop_box).float()
|
||||
near_crop_edge = torch.isclose(boxes, crop_box_torch[None, :], atol=atol, rtol=0)
|
||||
near_image_edge = torch.isclose(boxes, orig_box_torch[None, :], atol=atol, rtol=0)
|
||||
near_crop_edge = torch.logical_and(near_crop_edge, ~near_image_edge)
|
||||
return torch.any(near_crop_edge, dim=1)
|
||||
|
||||
|
||||
def box_xyxy_to_xywh(box_xyxy: torch.Tensor) -> torch.Tensor:
|
||||
box_xywh = deepcopy(box_xyxy)
|
||||
box_xywh[2] = box_xywh[2] - box_xywh[0]
|
||||
box_xywh[3] = box_xywh[3] - box_xywh[1]
|
||||
return box_xywh
|
||||
|
||||
|
||||
def batch_iterator(batch_size: int, *args) -> Generator[List[Any], None, None]:
|
||||
assert len(args) > 0 and all(
|
||||
len(a) == len(args[0]) for a in args
|
||||
), "Batched iteration must have inputs of all the same size."
|
||||
n_batches = len(args[0]) // batch_size + int(len(args[0]) % batch_size != 0)
|
||||
for b in range(n_batches):
|
||||
yield [arg[b * batch_size : (b + 1) * batch_size] for arg in args]
|
||||
|
||||
|
||||
def mask_to_rle_pytorch(tensor: torch.Tensor) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Encodes masks to an uncompressed RLE, in the format expected by
|
||||
pycoco tools.
|
||||
"""
|
||||
# Put in fortran order and flatten h,w
|
||||
b, h, w = tensor.shape
|
||||
tensor = tensor.permute(0, 2, 1).flatten(1)
|
||||
|
||||
# Compute change indices
|
||||
diff = tensor[:, 1:] ^ tensor[:, :-1]
|
||||
change_indices = diff.nonzero()
|
||||
|
||||
# Encode run length
|
||||
out = []
|
||||
for i in range(b):
|
||||
cur_idxs = change_indices[change_indices[:, 0] == i, 1]
|
||||
cur_idxs = torch.cat(
|
||||
[
|
||||
torch.tensor([0], dtype=cur_idxs.dtype, device=cur_idxs.device),
|
||||
cur_idxs + 1,
|
||||
torch.tensor([h * w], dtype=cur_idxs.dtype, device=cur_idxs.device),
|
||||
]
|
||||
)
|
||||
btw_idxs = cur_idxs[1:] - cur_idxs[:-1]
|
||||
counts = [] if tensor[i, 0] == 0 else [0]
|
||||
counts.extend(btw_idxs.detach().cpu().tolist())
|
||||
out.append({"size": [h, w], "counts": counts})
|
||||
return out
|
||||
|
||||
|
||||
def rle_to_mask(rle: Dict[str, Any]) -> np.ndarray:
|
||||
"""Compute a binary mask from an uncompressed RLE."""
|
||||
h, w = rle["size"]
|
||||
mask = np.empty(h * w, dtype=bool)
|
||||
idx = 0
|
||||
parity = False
|
||||
for count in rle["counts"]:
|
||||
mask[idx : idx + count] = parity
|
||||
idx += count
|
||||
parity ^= True
|
||||
mask = mask.reshape(w, h)
|
||||
return mask.transpose() # Put in C order
|
||||
|
||||
|
||||
def area_from_rle(rle: Dict[str, Any]) -> int:
|
||||
return sum(rle["counts"][1::2])
|
||||
|
||||
|
||||
def calculate_stability_score(
|
||||
masks: torch.Tensor, mask_threshold: float, threshold_offset: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Computes the stability score for a batch of masks. The stability
|
||||
score is the IoU between the binary masks obtained by thresholding
|
||||
the predicted mask logits at high and low values.
|
||||
"""
|
||||
# One mask is always contained inside the other.
|
||||
# Save memory by preventing unnecessary cast to torch.int64
|
||||
intersections = (
|
||||
(masks > (mask_threshold + threshold_offset))
|
||||
.sum(-1, dtype=torch.int16)
|
||||
.sum(-1, dtype=torch.int32)
|
||||
)
|
||||
unions = (
|
||||
(masks > (mask_threshold - threshold_offset))
|
||||
.sum(-1, dtype=torch.int16)
|
||||
.sum(-1, dtype=torch.int32)
|
||||
)
|
||||
return intersections / unions
|
||||
|
||||
|
||||
def build_point_grid(n_per_side: int) -> np.ndarray:
|
||||
"""Generates a 2D grid of points evenly spaced in [0,1]x[0,1]."""
|
||||
offset = 1 / (2 * n_per_side)
|
||||
points_one_side = np.linspace(offset, 1 - offset, n_per_side)
|
||||
points_x = np.tile(points_one_side[None, :], (n_per_side, 1))
|
||||
points_y = np.tile(points_one_side[:, None], (1, n_per_side))
|
||||
points = np.stack([points_x, points_y], axis=-1).reshape(-1, 2)
|
||||
return points
|
||||
|
||||
|
||||
def build_all_layer_point_grids(
|
||||
n_per_side: int, n_layers: int, scale_per_layer: int
|
||||
) -> List[np.ndarray]:
|
||||
"""Generates point grids for all crop layers."""
|
||||
points_by_layer = []
|
||||
for i in range(n_layers + 1):
|
||||
n_points = int(n_per_side / (scale_per_layer**i))
|
||||
points_by_layer.append(build_point_grid(n_points))
|
||||
return points_by_layer
|
||||
|
||||
|
||||
def generate_crop_boxes(
|
||||
im_size: Tuple[int, ...], n_layers: int, overlap_ratio: float
|
||||
) -> Tuple[List[List[int]], List[int]]:
|
||||
"""
|
||||
Generates a list of crop boxes of different sizes. Each layer
|
||||
has (2**i)**2 boxes for the ith layer.
|
||||
"""
|
||||
crop_boxes, layer_idxs = [], []
|
||||
im_h, im_w = im_size
|
||||
short_side = min(im_h, im_w)
|
||||
|
||||
# Original image
|
||||
crop_boxes.append([0, 0, im_w, im_h])
|
||||
layer_idxs.append(0)
|
||||
|
||||
def crop_len(orig_len, n_crops, overlap):
|
||||
return int(math.ceil((overlap * (n_crops - 1) + orig_len) / n_crops))
|
||||
|
||||
for i_layer in range(n_layers):
|
||||
n_crops_per_side = 2 ** (i_layer + 1)
|
||||
overlap = int(overlap_ratio * short_side * (2 / n_crops_per_side))
|
||||
|
||||
crop_w = crop_len(im_w, n_crops_per_side, overlap)
|
||||
crop_h = crop_len(im_h, n_crops_per_side, overlap)
|
||||
|
||||
crop_box_x0 = [int((crop_w - overlap) * i) for i in range(n_crops_per_side)]
|
||||
crop_box_y0 = [int((crop_h - overlap) * i) for i in range(n_crops_per_side)]
|
||||
|
||||
# Crops in XYWH format
|
||||
for x0, y0 in product(crop_box_x0, crop_box_y0):
|
||||
box = [x0, y0, min(x0 + crop_w, im_w), min(y0 + crop_h, im_h)]
|
||||
crop_boxes.append(box)
|
||||
layer_idxs.append(i_layer + 1)
|
||||
|
||||
return crop_boxes, layer_idxs
|
||||
|
||||
|
||||
def uncrop_boxes_xyxy(boxes: torch.Tensor, crop_box: List[int]) -> torch.Tensor:
|
||||
x0, y0, _, _ = crop_box
|
||||
offset = torch.tensor([[x0, y0, x0, y0]], device=boxes.device)
|
||||
# Check if boxes has a channel dimension
|
||||
if len(boxes.shape) == 3:
|
||||
offset = offset.unsqueeze(1)
|
||||
return boxes + offset
|
||||
|
||||
|
||||
def uncrop_points(points: torch.Tensor, crop_box: List[int]) -> torch.Tensor:
|
||||
x0, y0, _, _ = crop_box
|
||||
offset = torch.tensor([[x0, y0]], device=points.device)
|
||||
# Check if points has a channel dimension
|
||||
if len(points.shape) == 3:
|
||||
offset = offset.unsqueeze(1)
|
||||
return points + offset
|
||||
|
||||
|
||||
def uncrop_masks(
|
||||
masks: torch.Tensor, crop_box: List[int], orig_h: int, orig_w: int
|
||||
) -> torch.Tensor:
|
||||
x0, y0, x1, y1 = crop_box
|
||||
if x0 == 0 and y0 == 0 and x1 == orig_w and y1 == orig_h:
|
||||
return masks
|
||||
# Coordinate transform masks
|
||||
pad_x, pad_y = orig_w - (x1 - x0), orig_h - (y1 - y0)
|
||||
pad = (x0, pad_x - x0, y0, pad_y - y0)
|
||||
return torch.nn.functional.pad(masks, pad, value=0)
|
||||
|
||||
|
||||
def remove_small_regions(
|
||||
mask: np.ndarray, area_thresh: float, mode: str
|
||||
) -> Tuple[np.ndarray, bool]:
|
||||
"""
|
||||
Removes small disconnected regions and holes in a mask. Returns the
|
||||
mask and an indicator of if the mask has been modified.
|
||||
"""
|
||||
import cv2 # type: ignore
|
||||
|
||||
assert mode in ["holes", "islands"]
|
||||
correct_holes = mode == "holes"
|
||||
working_mask = (correct_holes ^ mask).astype(np.uint8)
|
||||
n_labels, regions, stats, _ = cv2.connectedComponentsWithStats(working_mask, 8)
|
||||
sizes = stats[:, -1][1:] # Row 0 is background label
|
||||
small_regions = [i + 1 for i, s in enumerate(sizes) if s < area_thresh]
|
||||
if len(small_regions) == 0:
|
||||
return mask, False
|
||||
fill_labels = [0] + small_regions
|
||||
if not correct_holes:
|
||||
fill_labels = [i for i in range(n_labels) if i not in fill_labels]
|
||||
# If every region is below threshold, keep largest
|
||||
if len(fill_labels) == 0:
|
||||
fill_labels = [int(np.argmax(sizes)) + 1]
|
||||
mask = np.isin(regions, fill_labels)
|
||||
return mask, True
|
||||
|
||||
|
||||
def coco_encode_rle(uncompressed_rle: Dict[str, Any]) -> Dict[str, Any]:
|
||||
from pycocotools import mask as mask_utils # type: ignore
|
||||
|
||||
h, w = uncompressed_rle["size"]
|
||||
rle = mask_utils.frPyObjects(uncompressed_rle, h, w)
|
||||
rle["counts"] = rle["counts"].decode("utf-8") # Necessary to serialize with json
|
||||
return rle
|
||||
|
||||
|
||||
def batched_mask_to_box(masks: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Calculates boxes in XYXY format around masks. Return [0,0,0,0] for
|
||||
an empty mask. For input shape C1xC2x...xHxW, the output shape is C1xC2x...x4.
|
||||
"""
|
||||
# torch.max below raises an error on empty inputs, just skip in this case
|
||||
if torch.numel(masks) == 0:
|
||||
return torch.zeros(*masks.shape[:-2], 4, device=masks.device)
|
||||
|
||||
# Normalize shape to CxHxW
|
||||
shape = masks.shape
|
||||
h, w = shape[-2:]
|
||||
if len(shape) > 2:
|
||||
masks = masks.flatten(0, -3)
|
||||
else:
|
||||
masks = masks.unsqueeze(0)
|
||||
|
||||
# Get top and bottom edges
|
||||
in_height, _ = torch.max(masks, dim=-1)
|
||||
in_height_coords = in_height * torch.arange(h, device=in_height.device)[None, :]
|
||||
bottom_edges, _ = torch.max(in_height_coords, dim=-1)
|
||||
in_height_coords = in_height_coords + h * (~in_height)
|
||||
top_edges, _ = torch.min(in_height_coords, dim=-1)
|
||||
|
||||
# Get left and right edges
|
||||
in_width, _ = torch.max(masks, dim=-2)
|
||||
in_width_coords = in_width * torch.arange(w, device=in_width.device)[None, :]
|
||||
right_edges, _ = torch.max(in_width_coords, dim=-1)
|
||||
in_width_coords = in_width_coords + w * (~in_width)
|
||||
left_edges, _ = torch.min(in_width_coords, dim=-1)
|
||||
|
||||
# If the mask is empty the right edge will be to the left of the left edge.
|
||||
# Replace these boxes with [0, 0, 0, 0]
|
||||
empty_filter = (right_edges < left_edges) | (bottom_edges < top_edges)
|
||||
out = torch.stack([left_edges, top_edges, right_edges, bottom_edges], dim=-1)
|
||||
out = out * (~empty_filter).unsqueeze(-1)
|
||||
|
||||
# Return to original shape
|
||||
if len(shape) > 2:
|
||||
out = out.reshape(*shape[:-2], 4)
|
||||
else:
|
||||
out = out[0]
|
||||
|
||||
return out
|
||||
@@ -0,0 +1,157 @@
|
||||
# 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.
|
||||
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.nn import functional as F
|
||||
|
||||
from ..modeling import Sam
|
||||
from .amg import calculate_stability_score
|
||||
|
||||
|
||||
class SamOnnxModel(nn.Module):
|
||||
"""
|
||||
This model should not be called directly, but is used in ONNX export.
|
||||
It combines the prompt encoder, mask decoder, and mask postprocessing of Sam,
|
||||
with some functions modified to enable model tracing. Also supports extra
|
||||
options controlling what information. See the ONNX export script for details.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: Sam,
|
||||
return_single_mask: bool,
|
||||
use_stability_score: bool = False,
|
||||
return_extra_metrics: bool = False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.mask_decoder = model.mask_decoder
|
||||
self.model = model
|
||||
self.img_size = model.image_encoder.img_size
|
||||
self.return_single_mask = return_single_mask
|
||||
self.use_stability_score = use_stability_score
|
||||
self.stability_score_offset = 1.0
|
||||
self.return_extra_metrics = return_extra_metrics
|
||||
|
||||
@staticmethod
|
||||
def resize_longest_image_size(
|
||||
input_image_size: torch.Tensor, longest_side: int
|
||||
) -> torch.Tensor:
|
||||
input_image_size = input_image_size.to(torch.float32)
|
||||
scale = longest_side / torch.max(input_image_size)
|
||||
transformed_size = scale * input_image_size
|
||||
transformed_size = torch.floor(transformed_size + 0.5).to(torch.int64)
|
||||
return transformed_size
|
||||
|
||||
def _embed_points(
|
||||
self, point_coords: torch.Tensor, point_labels: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
point_coords = point_coords + 0.5
|
||||
point_coords = point_coords / self.img_size
|
||||
point_embedding = self.model.prompt_encoder.pe_layer._pe_encoding(point_coords)
|
||||
point_labels = point_labels.unsqueeze(-1).expand_as(point_embedding)
|
||||
|
||||
point_embedding = point_embedding * (point_labels != -1)
|
||||
point_embedding = (
|
||||
point_embedding
|
||||
+ self.model.prompt_encoder.not_a_point_embed.weight * (point_labels == -1)
|
||||
)
|
||||
|
||||
for i in range(self.model.prompt_encoder.num_point_embeddings):
|
||||
point_embedding = (
|
||||
point_embedding
|
||||
+ self.model.prompt_encoder.point_embeddings[i].weight
|
||||
* (point_labels == i)
|
||||
)
|
||||
|
||||
return point_embedding
|
||||
|
||||
def _embed_masks(
|
||||
self, input_mask: torch.Tensor, has_mask_input: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
mask_embedding = has_mask_input * self.model.prompt_encoder.mask_downscaling(
|
||||
input_mask
|
||||
)
|
||||
mask_embedding = mask_embedding + (
|
||||
1 - has_mask_input
|
||||
) * self.model.prompt_encoder.no_mask_embed.weight.reshape(1, -1, 1, 1)
|
||||
return mask_embedding
|
||||
|
||||
def mask_postprocessing(
|
||||
self, masks: torch.Tensor, orig_im_size: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
masks = F.interpolate(
|
||||
masks,
|
||||
size=(self.img_size, self.img_size),
|
||||
mode="bilinear",
|
||||
align_corners=False,
|
||||
)
|
||||
|
||||
prepadded_size = self.resize_longest_image_size(orig_im_size, self.img_size).to(
|
||||
torch.int64
|
||||
)
|
||||
masks = masks[..., : prepadded_size[0], : prepadded_size[1]] # type: ignore
|
||||
|
||||
orig_im_size = orig_im_size.to(torch.int64)
|
||||
h, w = orig_im_size[0], orig_im_size[1]
|
||||
masks = F.interpolate(masks, size=(h, w), mode="bilinear", align_corners=False)
|
||||
return masks
|
||||
|
||||
def select_masks(
|
||||
self, masks: torch.Tensor, iou_preds: torch.Tensor, num_points: int
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
# Determine if we should return the multiclick mask or not from the number of points.
|
||||
# The reweighting is used to avoid control flow.
|
||||
score_reweight = torch.tensor(
|
||||
[[1000] + [0] * (self.model.mask_decoder.num_mask_tokens - 1)]
|
||||
).to(iou_preds.device)
|
||||
score = iou_preds + (num_points - 2.5) * score_reweight
|
||||
best_idx = torch.argmax(score, dim=1)
|
||||
masks = masks[torch.arange(masks.shape[0]), best_idx, :, :].unsqueeze(1)
|
||||
iou_preds = iou_preds[torch.arange(masks.shape[0]), best_idx].unsqueeze(1)
|
||||
|
||||
return masks, iou_preds
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
image_embeddings: torch.Tensor,
|
||||
point_coords: torch.Tensor,
|
||||
point_labels: torch.Tensor,
|
||||
mask_input: torch.Tensor,
|
||||
has_mask_input: torch.Tensor,
|
||||
orig_im_size: torch.Tensor,
|
||||
):
|
||||
sparse_embedding = self._embed_points(point_coords, point_labels)
|
||||
dense_embedding = self._embed_masks(mask_input, has_mask_input)
|
||||
|
||||
masks, scores = self.model.mask_decoder.predict_masks(
|
||||
image_embeddings=image_embeddings,
|
||||
image_pe=self.model.prompt_encoder.get_dense_pe(),
|
||||
sparse_prompt_embeddings=sparse_embedding,
|
||||
dense_prompt_embeddings=dense_embedding,
|
||||
)
|
||||
|
||||
if self.use_stability_score:
|
||||
scores = calculate_stability_score(
|
||||
masks, self.model.mask_threshold, self.stability_score_offset
|
||||
)
|
||||
|
||||
if self.return_single_mask:
|
||||
masks, scores = self.select_masks(masks, scores, point_coords.shape[1])
|
||||
|
||||
upscaled_masks = self.mask_postprocessing(masks, orig_im_size)
|
||||
|
||||
if self.return_extra_metrics:
|
||||
stability_scores = calculate_stability_score(
|
||||
upscaled_masks, self.model.mask_threshold, self.stability_score_offset
|
||||
)
|
||||
areas = (upscaled_masks > self.model.mask_threshold).sum(-1).sum(-1)
|
||||
return upscaled_masks, scores, stability_scores, areas, masks
|
||||
|
||||
return upscaled_masks, scores, masks
|
||||
@@ -0,0 +1,113 @@
|
||||
# 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.
|
||||
|
||||
from copy import deepcopy
|
||||
from typing import Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch.nn import functional as F
|
||||
from torchvision.transforms.functional import resize # type: ignore
|
||||
from torchvision.transforms.functional import to_pil_image
|
||||
|
||||
|
||||
class ResizeLongestSide:
|
||||
"""
|
||||
Resizes images to the longest side 'target_length', as well as provides
|
||||
methods for resizing coordinates and boxes. Provides methods for
|
||||
transforming both numpy array and batched torch tensors.
|
||||
"""
|
||||
|
||||
def __init__(self, target_length: int) -> None:
|
||||
self.target_length = target_length
|
||||
|
||||
def apply_image(self, image: np.ndarray) -> np.ndarray:
|
||||
"""
|
||||
Expects a numpy array with shape HxWxC in uint8 format.
|
||||
"""
|
||||
target_size = self.get_preprocess_shape(
|
||||
image.shape[0], image.shape[1], self.target_length
|
||||
)
|
||||
return np.array(resize(to_pil_image(image), target_size))
|
||||
|
||||
def apply_coords(
|
||||
self, coords: np.ndarray, original_size: Tuple[int, ...]
|
||||
) -> np.ndarray:
|
||||
"""
|
||||
Expects a numpy array of length 2 in the final dimension. Requires the
|
||||
original image size in (H, W) format.
|
||||
"""
|
||||
old_h, old_w = original_size
|
||||
new_h, new_w = self.get_preprocess_shape(
|
||||
original_size[0], original_size[1], self.target_length
|
||||
)
|
||||
coords = deepcopy(coords).astype(float)
|
||||
coords[..., 0] = coords[..., 0] * (new_w / old_w)
|
||||
coords[..., 1] = coords[..., 1] * (new_h / old_h)
|
||||
return coords
|
||||
|
||||
def apply_boxes(
|
||||
self, boxes: np.ndarray, original_size: Tuple[int, ...]
|
||||
) -> np.ndarray:
|
||||
"""
|
||||
Expects a numpy array shape Bx4. Requires the original image size
|
||||
in (H, W) format.
|
||||
"""
|
||||
boxes = self.apply_coords(boxes.reshape(-1, 2, 2), original_size)
|
||||
return boxes.reshape(-1, 4)
|
||||
|
||||
def apply_image_torch(self, image: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Expects batched images with shape BxCxHxW and float format. This
|
||||
transformation may not exactly match apply_image. apply_image is
|
||||
the transformation expected by the model.
|
||||
"""
|
||||
# Expects an image in BCHW format. May not exactly match apply_image.
|
||||
target_size = self.get_preprocess_shape(
|
||||
image.shape[0], image.shape[1], self.target_length
|
||||
)
|
||||
return F.interpolate(
|
||||
image, target_size, mode="bilinear", align_corners=False, antialias=True
|
||||
)
|
||||
|
||||
def apply_coords_torch(
|
||||
self, coords: torch.Tensor, original_size: Tuple[int, ...]
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Expects a torch tensor with length 2 in the last dimension. Requires the
|
||||
original image size in (H, W) format.
|
||||
"""
|
||||
old_h, old_w = original_size
|
||||
new_h, new_w = self.get_preprocess_shape(
|
||||
original_size[0], original_size[1], self.target_length
|
||||
)
|
||||
coords = deepcopy(coords).to(torch.float)
|
||||
coords[..., 0] = coords[..., 0] * (new_w / old_w)
|
||||
coords[..., 1] = coords[..., 1] * (new_h / old_h)
|
||||
return coords
|
||||
|
||||
def apply_boxes_torch(
|
||||
self, boxes: torch.Tensor, original_size: Tuple[int, ...]
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Expects a torch tensor with shape Bx4. Requires the original image
|
||||
size in (H, W) format.
|
||||
"""
|
||||
boxes = self.apply_coords_torch(boxes.reshape(-1, 2, 2), original_size)
|
||||
return boxes.reshape(-1, 4)
|
||||
|
||||
@staticmethod
|
||||
def get_preprocess_shape(
|
||||
oldh: int, oldw: int, long_side_length: int
|
||||
) -> Tuple[int, int]:
|
||||
"""
|
||||
Compute the output size given input size and target long side length.
|
||||
"""
|
||||
scale = long_side_length * 1.0 / max(oldh, oldw)
|
||||
newh, neww = oldh * scale, oldw * scale
|
||||
neww = int(neww + 0.5)
|
||||
newh = int(newh + 0.5)
|
||||
return (newh, neww)
|
||||
@@ -0,0 +1,10 @@
|
||||
# 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.
|
||||
|
||||
|
||||
from hydra import initialize_config_module
|
||||
|
||||
initialize_config_module("model/segment_anything_2/sam2_configs", version_base="1.2")
|
||||
@@ -0,0 +1,434 @@
|
||||
# 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.
|
||||
|
||||
# Adapted from https://github.com/facebookresearch/segment-anything/blob/main/segment_anything/automatic_mask_generator.py
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from torchvision.ops.boxes import batched_nms, box_area # type: ignore
|
||||
|
||||
from model.segment_anything_2.sam2.modeling.sam2_base import SAM2Base
|
||||
from model.segment_anything_2.sam2.sam2_image_predictor import SAM2ImagePredictor
|
||||
from model.segment_anything_2.sam2.utils.amg import (
|
||||
area_from_rle,
|
||||
batch_iterator,
|
||||
batched_mask_to_box,
|
||||
box_xyxy_to_xywh,
|
||||
build_all_layer_point_grids,
|
||||
calculate_stability_score,
|
||||
coco_encode_rle,
|
||||
generate_crop_boxes,
|
||||
is_box_near_crop_edge,
|
||||
mask_to_rle_pytorch,
|
||||
MaskData,
|
||||
remove_small_regions,
|
||||
rle_to_mask,
|
||||
uncrop_boxes_xyxy,
|
||||
uncrop_masks,
|
||||
uncrop_points,
|
||||
)
|
||||
|
||||
|
||||
class SAM2AutomaticMaskGenerator:
|
||||
def __init__(
|
||||
self,
|
||||
model: SAM2Base,
|
||||
points_per_side: Optional[int] = 32,
|
||||
points_per_batch: int = 64,
|
||||
pred_iou_thresh: float = 0.8,
|
||||
stability_score_thresh: float = 0.95,
|
||||
stability_score_offset: float = 1.0,
|
||||
mask_threshold: float = 0.0,
|
||||
box_nms_thresh: float = 0.7,
|
||||
crop_n_layers: int = 0,
|
||||
crop_nms_thresh: float = 0.7,
|
||||
crop_overlap_ratio: float = 512 / 1500,
|
||||
crop_n_points_downscale_factor: int = 1,
|
||||
point_grids: Optional[List[np.ndarray]] = None,
|
||||
min_mask_region_area: int = 0,
|
||||
output_mode: str = "binary_mask",
|
||||
use_m2m: bool = False,
|
||||
multimask_output: bool = True,
|
||||
) -> None:
|
||||
"""
|
||||
Using a SAM 2 model, generates masks for the entire image.
|
||||
Generates a grid of point prompts over the image, then filters
|
||||
low quality and duplicate masks. The default settings are chosen
|
||||
for SAM 2 with a HieraL backbone.
|
||||
|
||||
Arguments:
|
||||
model (Sam): The SAM 2 model to use for mask prediction.
|
||||
points_per_side (int or None): The number of points to be sampled
|
||||
along one side of the image. The total number of points is
|
||||
points_per_side**2. If None, 'point_grids' must provide explicit
|
||||
point sampling.
|
||||
points_per_batch (int): Sets the number of points run simultaneously
|
||||
by the model. Higher numbers may be faster but use more GPU memory.
|
||||
pred_iou_thresh (float): A filtering threshold in [0,1], using the
|
||||
model's predicted mask quality.
|
||||
stability_score_thresh (float): A filtering threshold in [0,1], using
|
||||
the stability of the mask under changes to the cutoff used to binarize
|
||||
the model's mask predictions.
|
||||
stability_score_offset (float): The amount to shift the cutoff when
|
||||
calculated the stability score.
|
||||
mask_threshold (float): Threshold for binarizing the mask logits
|
||||
box_nms_thresh (float): The box IoU cutoff used by non-maximal
|
||||
suppression to filter duplicate masks.
|
||||
crop_n_layers (int): If >0, mask prediction will be run again on
|
||||
crops of the image. Sets the number of layers to run, where each
|
||||
layer has 2**i_layer number of image crops.
|
||||
crop_nms_thresh (float): The box IoU cutoff used by non-maximal
|
||||
suppression to filter duplicate masks between different crops.
|
||||
crop_overlap_ratio (float): Sets the degree to which crops overlap.
|
||||
In the first crop layer, crops will overlap by this fraction of
|
||||
the image length. Later layers with more crops scale down this overlap.
|
||||
crop_n_points_downscale_factor (int): The number of points-per-side
|
||||
sampled in layer n is scaled down by crop_n_points_downscale_factor**n.
|
||||
point_grids (list(np.ndarray) or None): A list over explicit grids
|
||||
of points used for sampling, normalized to [0,1]. The nth grid in the
|
||||
list is used in the nth crop layer. Exclusive with points_per_side.
|
||||
min_mask_region_area (int): If >0, postprocessing will be applied
|
||||
to remove disconnected regions and holes in masks with area smaller
|
||||
than min_mask_region_area. Requires opencv.
|
||||
output_mode (str): The form masks are returned in. Can be 'binary_mask',
|
||||
'uncompressed_rle', or 'coco_rle'. 'coco_rle' requires pycocotools.
|
||||
For large resolutions, 'binary_mask' may consume large amounts of
|
||||
memory.
|
||||
use_m2m (bool): Whether to add a one step refinement using previous mask predictions.
|
||||
multimask_output (bool): Whether to output multimask at each point of the grid.
|
||||
"""
|
||||
|
||||
assert (points_per_side is None) != (
|
||||
point_grids is None
|
||||
), "Exactly one of points_per_side or point_grid must be provided."
|
||||
if points_per_side is not None:
|
||||
self.point_grids = build_all_layer_point_grids(
|
||||
points_per_side,
|
||||
crop_n_layers,
|
||||
crop_n_points_downscale_factor,
|
||||
)
|
||||
elif point_grids is not None:
|
||||
self.point_grids = point_grids
|
||||
else:
|
||||
raise ValueError("Can't have both points_per_side and point_grid be None.")
|
||||
|
||||
assert output_mode in [
|
||||
"binary_mask",
|
||||
"uncompressed_rle",
|
||||
"coco_rle",
|
||||
], f"Unknown output_mode {output_mode}."
|
||||
if output_mode == "coco_rle":
|
||||
try:
|
||||
from pycocotools import mask as mask_utils # type: ignore # noqa: F401
|
||||
except ImportError as e:
|
||||
print("Please install pycocotools")
|
||||
raise e
|
||||
|
||||
self.predictor = SAM2ImagePredictor(
|
||||
model,
|
||||
max_hole_area=min_mask_region_area,
|
||||
max_sprinkle_area=min_mask_region_area,
|
||||
)
|
||||
self.points_per_batch = points_per_batch
|
||||
self.pred_iou_thresh = pred_iou_thresh
|
||||
self.stability_score_thresh = stability_score_thresh
|
||||
self.stability_score_offset = stability_score_offset
|
||||
self.mask_threshold = mask_threshold
|
||||
self.box_nms_thresh = box_nms_thresh
|
||||
self.crop_n_layers = crop_n_layers
|
||||
self.crop_nms_thresh = crop_nms_thresh
|
||||
self.crop_overlap_ratio = crop_overlap_ratio
|
||||
self.crop_n_points_downscale_factor = crop_n_points_downscale_factor
|
||||
self.min_mask_region_area = min_mask_region_area
|
||||
self.output_mode = output_mode
|
||||
self.use_m2m = use_m2m
|
||||
self.multimask_output = multimask_output
|
||||
|
||||
@torch.no_grad()
|
||||
def generate(self, image: np.ndarray) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Generates masks for the given image.
|
||||
|
||||
Arguments:
|
||||
image (np.ndarray): The image to generate masks for, in HWC uint8 format.
|
||||
|
||||
Returns:
|
||||
list(dict(str, any)): A list over records for masks. Each record is
|
||||
a dict containing the following keys:
|
||||
segmentation (dict(str, any) or np.ndarray): The mask. If
|
||||
output_mode='binary_mask', is an array of shape HW. Otherwise,
|
||||
is a dictionary containing the RLE.
|
||||
bbox (list(float)): The box around the mask, in XYWH format.
|
||||
area (int): The area in pixels of the mask.
|
||||
predicted_iou (float): The model's own prediction of the mask's
|
||||
quality. This is filtered by the pred_iou_thresh parameter.
|
||||
point_coords (list(list(float))): The point coordinates input
|
||||
to the model to generate this mask.
|
||||
stability_score (float): A measure of the mask's quality. This
|
||||
is filtered on using the stability_score_thresh parameter.
|
||||
crop_box (list(float)): The crop of the image used to generate
|
||||
the mask, given in XYWH format.
|
||||
"""
|
||||
|
||||
# Generate masks
|
||||
mask_data = self._generate_masks(image)
|
||||
|
||||
# Encode masks
|
||||
if self.output_mode == "coco_rle":
|
||||
mask_data["segmentations"] = [
|
||||
coco_encode_rle(rle) for rle in mask_data["rles"]
|
||||
]
|
||||
elif self.output_mode == "binary_mask":
|
||||
mask_data["segmentations"] = [rle_to_mask(rle) for rle in mask_data["rles"]]
|
||||
else:
|
||||
mask_data["segmentations"] = mask_data["rles"]
|
||||
|
||||
# Write mask records
|
||||
curr_anns = []
|
||||
for idx in range(len(mask_data["segmentations"])):
|
||||
ann = {
|
||||
"segmentation": mask_data["segmentations"][idx],
|
||||
"area": area_from_rle(mask_data["rles"][idx]),
|
||||
"bbox": box_xyxy_to_xywh(mask_data["boxes"][idx]).tolist(),
|
||||
"predicted_iou": mask_data["iou_preds"][idx].item(),
|
||||
"point_coords": [mask_data["points"][idx].tolist()],
|
||||
"stability_score": mask_data["stability_score"][idx].item(),
|
||||
"crop_box": box_xyxy_to_xywh(mask_data["crop_boxes"][idx]).tolist(),
|
||||
}
|
||||
curr_anns.append(ann)
|
||||
|
||||
return curr_anns
|
||||
|
||||
def _generate_masks(self, image: np.ndarray) -> MaskData:
|
||||
orig_size = image.shape[:2]
|
||||
crop_boxes, layer_idxs = generate_crop_boxes(
|
||||
orig_size, self.crop_n_layers, self.crop_overlap_ratio
|
||||
)
|
||||
|
||||
# Iterate over image crops
|
||||
data = MaskData()
|
||||
for crop_box, layer_idx in zip(crop_boxes, layer_idxs):
|
||||
crop_data = self._process_crop(image, crop_box, layer_idx, orig_size)
|
||||
data.cat(crop_data)
|
||||
|
||||
# Remove duplicate masks between crops
|
||||
if len(crop_boxes) > 1:
|
||||
# Prefer masks from smaller crops
|
||||
scores = 1 / box_area(data["crop_boxes"])
|
||||
scores = scores.to(data["boxes"].device)
|
||||
keep_by_nms = batched_nms(
|
||||
data["boxes"].float(),
|
||||
scores,
|
||||
torch.zeros_like(data["boxes"][:, 0]), # categories
|
||||
iou_threshold=self.crop_nms_thresh,
|
||||
)
|
||||
data.filter(keep_by_nms)
|
||||
data.to_numpy()
|
||||
return data
|
||||
|
||||
def _process_crop(
|
||||
self,
|
||||
image: np.ndarray,
|
||||
crop_box: List[int],
|
||||
crop_layer_idx: int,
|
||||
orig_size: Tuple[int, ...],
|
||||
) -> MaskData:
|
||||
# Crop the image and calculate embeddings
|
||||
x0, y0, x1, y1 = crop_box
|
||||
cropped_im = image[y0:y1, x0:x1, :]
|
||||
cropped_im_size = cropped_im.shape[:2]
|
||||
self.predictor.set_image(cropped_im)
|
||||
|
||||
# Get points for this crop
|
||||
points_scale = np.array(cropped_im_size)[None, ::-1]
|
||||
points_for_image = self.point_grids[crop_layer_idx] * points_scale
|
||||
|
||||
# Generate masks for this crop in batches
|
||||
data = MaskData()
|
||||
for (points,) in batch_iterator(self.points_per_batch, points_for_image):
|
||||
batch_data = self._process_batch(
|
||||
points, cropped_im_size, crop_box, orig_size, normalize=True
|
||||
)
|
||||
data.cat(batch_data)
|
||||
del batch_data
|
||||
self.predictor.reset_predictor()
|
||||
|
||||
# Remove duplicates within this crop.
|
||||
keep_by_nms = batched_nms(
|
||||
data["boxes"].float(),
|
||||
data["iou_preds"],
|
||||
torch.zeros_like(data["boxes"][:, 0]), # categories
|
||||
iou_threshold=self.box_nms_thresh,
|
||||
)
|
||||
data.filter(keep_by_nms)
|
||||
|
||||
# Return to the original image frame
|
||||
data["boxes"] = uncrop_boxes_xyxy(data["boxes"], crop_box)
|
||||
data["points"] = uncrop_points(data["points"], crop_box)
|
||||
data["crop_boxes"] = torch.tensor([crop_box for _ in range(len(data["rles"]))])
|
||||
|
||||
return data
|
||||
|
||||
def _process_batch(
|
||||
self,
|
||||
points: np.ndarray,
|
||||
im_size: Tuple[int, ...],
|
||||
crop_box: List[int],
|
||||
orig_size: Tuple[int, ...],
|
||||
normalize=False,
|
||||
) -> MaskData:
|
||||
orig_h, orig_w = orig_size
|
||||
|
||||
# Run model on this batch
|
||||
points = torch.as_tensor(points, device=self.predictor.device)
|
||||
in_points = self.predictor._transforms.transform_coords(
|
||||
points, normalize=normalize, orig_hw=im_size
|
||||
)
|
||||
in_labels = torch.ones(
|
||||
in_points.shape[0], dtype=torch.int, device=in_points.device
|
||||
)
|
||||
masks, iou_preds, low_res_masks = self.predictor._predict(
|
||||
in_points[:, None, :],
|
||||
in_labels[:, None],
|
||||
multimask_output=self.multimask_output,
|
||||
return_logits=True,
|
||||
)
|
||||
|
||||
# Serialize predictions and store in MaskData
|
||||
data = MaskData(
|
||||
masks=masks.flatten(0, 1),
|
||||
iou_preds=iou_preds.flatten(0, 1),
|
||||
points=points.repeat_interleave(masks.shape[1], dim=0),
|
||||
low_res_masks=low_res_masks.flatten(0, 1),
|
||||
)
|
||||
del masks
|
||||
|
||||
if not self.use_m2m:
|
||||
# Filter by predicted IoU
|
||||
if self.pred_iou_thresh > 0.0:
|
||||
keep_mask = data["iou_preds"] > self.pred_iou_thresh
|
||||
data.filter(keep_mask)
|
||||
|
||||
# Calculate and filter by stability score
|
||||
data["stability_score"] = calculate_stability_score(
|
||||
data["masks"], self.mask_threshold, self.stability_score_offset
|
||||
)
|
||||
if self.stability_score_thresh > 0.0:
|
||||
keep_mask = data["stability_score"] >= self.stability_score_thresh
|
||||
data.filter(keep_mask)
|
||||
else:
|
||||
# One step refinement using previous mask predictions
|
||||
in_points = self.predictor._transforms.transform_coords(
|
||||
data["points"], normalize=normalize, orig_hw=im_size
|
||||
)
|
||||
labels = torch.ones(
|
||||
in_points.shape[0], dtype=torch.int, device=in_points.device
|
||||
)
|
||||
masks, ious = self.refine_with_m2m(
|
||||
in_points, labels, data["low_res_masks"], self.points_per_batch
|
||||
)
|
||||
data["masks"] = masks.squeeze(1)
|
||||
data["iou_preds"] = ious.squeeze(1)
|
||||
|
||||
if self.pred_iou_thresh > 0.0:
|
||||
keep_mask = data["iou_preds"] > self.pred_iou_thresh
|
||||
data.filter(keep_mask)
|
||||
|
||||
data["stability_score"] = calculate_stability_score(
|
||||
data["masks"], self.mask_threshold, self.stability_score_offset
|
||||
)
|
||||
if self.stability_score_thresh > 0.0:
|
||||
keep_mask = data["stability_score"] >= self.stability_score_thresh
|
||||
data.filter(keep_mask)
|
||||
|
||||
# Threshold masks and calculate boxes
|
||||
data["masks"] = data["masks"] > self.mask_threshold
|
||||
data["boxes"] = batched_mask_to_box(data["masks"])
|
||||
|
||||
# Filter boxes that touch crop boundaries
|
||||
keep_mask = ~is_box_near_crop_edge(
|
||||
data["boxes"], crop_box, [0, 0, orig_w, orig_h]
|
||||
)
|
||||
if not torch.all(keep_mask):
|
||||
data.filter(keep_mask)
|
||||
|
||||
# Compress to RLE
|
||||
data["masks"] = uncrop_masks(data["masks"], crop_box, orig_h, orig_w)
|
||||
data["rles"] = mask_to_rle_pytorch(data["masks"])
|
||||
del data["masks"]
|
||||
|
||||
return data
|
||||
|
||||
@staticmethod
|
||||
def postprocess_small_regions(
|
||||
mask_data: MaskData, min_area: int, nms_thresh: float
|
||||
) -> MaskData:
|
||||
"""
|
||||
Removes small disconnected regions and holes in masks, then reruns
|
||||
box NMS to remove any new duplicates.
|
||||
|
||||
Edits mask_data in place.
|
||||
|
||||
Requires open-cv as a dependency.
|
||||
"""
|
||||
if len(mask_data["rles"]) == 0:
|
||||
return mask_data
|
||||
|
||||
# Filter small disconnected regions and holes
|
||||
new_masks = []
|
||||
scores = []
|
||||
for rle in mask_data["rles"]:
|
||||
mask = rle_to_mask(rle)
|
||||
|
||||
mask, changed = remove_small_regions(mask, min_area, mode="holes")
|
||||
unchanged = not changed
|
||||
mask, changed = remove_small_regions(mask, min_area, mode="islands")
|
||||
unchanged = unchanged and not changed
|
||||
|
||||
new_masks.append(torch.as_tensor(mask).unsqueeze(0))
|
||||
# Give score=0 to changed masks and score=1 to unchanged masks
|
||||
# so NMS will prefer ones that didn't need postprocessing
|
||||
scores.append(float(unchanged))
|
||||
|
||||
# Recalculate boxes and remove any new duplicates
|
||||
masks = torch.cat(new_masks, dim=0)
|
||||
boxes = batched_mask_to_box(masks)
|
||||
keep_by_nms = batched_nms(
|
||||
boxes.float(),
|
||||
torch.as_tensor(scores),
|
||||
torch.zeros_like(boxes[:, 0]), # categories
|
||||
iou_threshold=nms_thresh,
|
||||
)
|
||||
|
||||
# Only recalculate RLEs for masks that have changed
|
||||
for i_mask in keep_by_nms:
|
||||
if scores[i_mask] == 0.0:
|
||||
mask_torch = masks[i_mask].unsqueeze(0)
|
||||
mask_data["rles"][i_mask] = mask_to_rle_pytorch(mask_torch)[0]
|
||||
mask_data["boxes"][i_mask] = boxes[i_mask] # update res directly
|
||||
mask_data.filter(keep_by_nms)
|
||||
|
||||
return mask_data
|
||||
|
||||
def refine_with_m2m(self, points, point_labels, low_res_masks, points_per_batch):
|
||||
new_masks = []
|
||||
new_iou_preds = []
|
||||
|
||||
for cur_points, cur_point_labels, low_res_mask in batch_iterator(
|
||||
points_per_batch, points, point_labels, low_res_masks
|
||||
):
|
||||
best_masks, best_iou_preds, _ = self.predictor._predict(
|
||||
cur_points[:, None, :],
|
||||
cur_point_labels[:, None],
|
||||
mask_input=low_res_mask[:, None, :],
|
||||
multimask_output=False,
|
||||
return_logits=True,
|
||||
)
|
||||
new_masks.append(best_masks)
|
||||
new_iou_preds.append(best_iou_preds)
|
||||
masks = torch.cat(new_masks, dim=0)
|
||||
return masks, torch.cat(new_iou_preds, dim=0)
|
||||
@@ -0,0 +1,90 @@
|
||||
# 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 logging
|
||||
|
||||
import torch
|
||||
from hydra import compose
|
||||
from hydra.utils import instantiate
|
||||
from omegaconf import OmegaConf
|
||||
|
||||
def build_sam2(
|
||||
config_file,
|
||||
ckpt_path=None,
|
||||
device="cuda",
|
||||
mode="eval",
|
||||
hydra_overrides_extra=[],
|
||||
apply_postprocessing=True,
|
||||
):
|
||||
|
||||
if apply_postprocessing:
|
||||
hydra_overrides_extra = hydra_overrides_extra.copy()
|
||||
hydra_overrides_extra += [
|
||||
# dynamically fall back to multi-mask if the single mask is not stable
|
||||
"++model.sam_mask_decoder_extra_args.dynamic_multimask_via_stability=true",
|
||||
"++model.sam_mask_decoder_extra_args.dynamic_multimask_stability_delta=0.05",
|
||||
"++model.sam_mask_decoder_extra_args.dynamic_multimask_stability_thresh=0.98",
|
||||
]
|
||||
# Read config and init model
|
||||
cfg = compose(config_name=config_file, overrides=hydra_overrides_extra)
|
||||
OmegaConf.resolve(cfg)
|
||||
model = instantiate(cfg.model, _recursive_=True)
|
||||
_load_checkpoint(model, ckpt_path)
|
||||
if device:
|
||||
model = model.to(device)
|
||||
if mode == "eval":
|
||||
model.eval()
|
||||
return model
|
||||
|
||||
|
||||
def build_sam2_video_predictor(
|
||||
config_file,
|
||||
ckpt_path=None,
|
||||
device="cuda",
|
||||
mode="eval",
|
||||
hydra_overrides_extra=[],
|
||||
apply_postprocessing=True,
|
||||
):
|
||||
hydra_overrides = [
|
||||
"++model._target_=model.segment_anything_2.sam2.sam2_video_predictor.SAM2VideoPredictor",
|
||||
]
|
||||
if apply_postprocessing:
|
||||
hydra_overrides_extra = hydra_overrides_extra.copy()
|
||||
hydra_overrides_extra += [
|
||||
# dynamically fall back to multi-mask if the single mask is not stable
|
||||
"++model.sam_mask_decoder_extra_args.dynamic_multimask_via_stability=true",
|
||||
"++model.sam_mask_decoder_extra_args.dynamic_multimask_stability_delta=0.05",
|
||||
"++model.sam_mask_decoder_extra_args.dynamic_multimask_stability_thresh=0.98",
|
||||
# the sigmoid mask logits on interacted frames with clicks in the memory encoder so that the encoded masks are exactly as what users see from clicking
|
||||
"++model.binarize_mask_from_pts_for_mem_enc=true",
|
||||
# fill small holes in the low-res masks up to `fill_hole_area` (before resizing them to the original video resolution)
|
||||
"++model.fill_hole_area=8",
|
||||
]
|
||||
hydra_overrides.extend(hydra_overrides_extra)
|
||||
|
||||
# Read config and init model
|
||||
cfg = compose(config_name=config_file, overrides=hydra_overrides)
|
||||
OmegaConf.resolve(cfg)
|
||||
model = instantiate(cfg.model, _recursive_=True)
|
||||
_load_checkpoint(model, ckpt_path)
|
||||
if device:
|
||||
model = model.to(device)
|
||||
if mode == "eval":
|
||||
model.eval()
|
||||
return model
|
||||
|
||||
|
||||
def _load_checkpoint(model, ckpt_path):
|
||||
if ckpt_path is not None:
|
||||
sd = torch.load(ckpt_path, map_location="cpu")["model"]
|
||||
missing_keys, unexpected_keys = model.load_state_dict(sd)
|
||||
if missing_keys:
|
||||
logging.error(missing_keys)
|
||||
raise RuntimeError()
|
||||
if unexpected_keys:
|
||||
logging.error(unexpected_keys)
|
||||
raise RuntimeError()
|
||||
logging.info("Loaded checkpoint sucessfully")
|
||||
@@ -0,0 +1,289 @@
|
||||
// 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.
|
||||
|
||||
// adapted from https://github.com/zsef123/Connected_components_PyTorch
|
||||
// with license found in the LICENSE_cctorch file in the root directory.
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <cuda.h>
|
||||
#include <cuda_runtime.h>
|
||||
#include <torch/extension.h>
|
||||
#include <torch/script.h>
|
||||
#include <vector>
|
||||
|
||||
// 2d
|
||||
#define BLOCK_ROWS 16
|
||||
#define BLOCK_COLS 16
|
||||
|
||||
namespace cc2d {
|
||||
|
||||
template <typename T>
|
||||
__device__ __forceinline__ unsigned char hasBit(T bitmap, unsigned char pos) {
|
||||
return (bitmap >> pos) & 1;
|
||||
}
|
||||
|
||||
__device__ int32_t find(const int32_t* s_buf, int32_t n) {
|
||||
while (s_buf[n] != n)
|
||||
n = s_buf[n];
|
||||
return n;
|
||||
}
|
||||
|
||||
__device__ int32_t find_n_compress(int32_t* s_buf, int32_t n) {
|
||||
const int32_t id = n;
|
||||
while (s_buf[n] != n) {
|
||||
n = s_buf[n];
|
||||
s_buf[id] = n;
|
||||
}
|
||||
return n;
|
||||
}
|
||||
|
||||
__device__ void union_(int32_t* s_buf, int32_t a, int32_t b) {
|
||||
bool done;
|
||||
do {
|
||||
a = find(s_buf, a);
|
||||
b = find(s_buf, b);
|
||||
|
||||
if (a < b) {
|
||||
int32_t old = atomicMin(s_buf + b, a);
|
||||
done = (old == b);
|
||||
b = old;
|
||||
} else if (b < a) {
|
||||
int32_t old = atomicMin(s_buf + a, b);
|
||||
done = (old == a);
|
||||
a = old;
|
||||
} else
|
||||
done = true;
|
||||
|
||||
} while (!done);
|
||||
}
|
||||
|
||||
__global__ void
|
||||
init_labeling(int32_t* label, const uint32_t W, const uint32_t H) {
|
||||
const uint32_t row = (blockIdx.y * blockDim.y + threadIdx.y) * 2;
|
||||
const uint32_t col = (blockIdx.x * blockDim.x + threadIdx.x) * 2;
|
||||
const uint32_t idx = row * W + col;
|
||||
|
||||
if (row < H && col < W)
|
||||
label[idx] = idx;
|
||||
}
|
||||
|
||||
__global__ void
|
||||
merge(uint8_t* img, int32_t* label, const uint32_t W, const uint32_t H) {
|
||||
const uint32_t row = (blockIdx.y * blockDim.y + threadIdx.y) * 2;
|
||||
const uint32_t col = (blockIdx.x * blockDim.x + threadIdx.x) * 2;
|
||||
const uint32_t idx = row * W + col;
|
||||
|
||||
if (row >= H || col >= W)
|
||||
return;
|
||||
|
||||
uint32_t P = 0;
|
||||
|
||||
if (img[idx])
|
||||
P |= 0x777;
|
||||
if (row + 1 < H && img[idx + W])
|
||||
P |= 0x777 << 4;
|
||||
if (col + 1 < W && img[idx + 1])
|
||||
P |= 0x777 << 1;
|
||||
|
||||
if (col == 0)
|
||||
P &= 0xEEEE;
|
||||
if (col + 1 >= W)
|
||||
P &= 0x3333;
|
||||
else if (col + 2 >= W)
|
||||
P &= 0x7777;
|
||||
|
||||
if (row == 0)
|
||||
P &= 0xFFF0;
|
||||
if (row + 1 >= H)
|
||||
P &= 0xFF;
|
||||
|
||||
if (P > 0) {
|
||||
// If need check about top-left pixel(if flag the first bit) and hit the
|
||||
// top-left pixel
|
||||
if (hasBit(P, 0) && img[idx - W - 1]) {
|
||||
union_(label, idx, idx - 2 * W - 2); // top left block
|
||||
}
|
||||
|
||||
if ((hasBit(P, 1) && img[idx - W]) || (hasBit(P, 2) && img[idx - W + 1]))
|
||||
union_(label, idx, idx - 2 * W); // top bottom block
|
||||
|
||||
if (hasBit(P, 3) && img[idx + 2 - W])
|
||||
union_(label, idx, idx - 2 * W + 2); // top right block
|
||||
|
||||
if ((hasBit(P, 4) && img[idx - 1]) || (hasBit(P, 8) && img[idx + W - 1]))
|
||||
union_(label, idx, idx - 2); // just left block
|
||||
}
|
||||
}
|
||||
|
||||
__global__ void compression(int32_t* label, const int32_t W, const int32_t H) {
|
||||
const uint32_t row = (blockIdx.y * blockDim.y + threadIdx.y) * 2;
|
||||
const uint32_t col = (blockIdx.x * blockDim.x + threadIdx.x) * 2;
|
||||
const uint32_t idx = row * W + col;
|
||||
|
||||
if (row < H && col < W)
|
||||
find_n_compress(label, idx);
|
||||
}
|
||||
|
||||
__global__ void final_labeling(
|
||||
const uint8_t* img,
|
||||
int32_t* label,
|
||||
const int32_t W,
|
||||
const int32_t H) {
|
||||
const uint32_t row = (blockIdx.y * blockDim.y + threadIdx.y) * 2;
|
||||
const uint32_t col = (blockIdx.x * blockDim.x + threadIdx.x) * 2;
|
||||
const uint32_t idx = row * W + col;
|
||||
|
||||
if (row >= H || col >= W)
|
||||
return;
|
||||
|
||||
int32_t y = label[idx] + 1;
|
||||
|
||||
if (img[idx])
|
||||
label[idx] = y;
|
||||
else
|
||||
label[idx] = 0;
|
||||
|
||||
if (col + 1 < W) {
|
||||
if (img[idx + 1])
|
||||
label[idx + 1] = y;
|
||||
else
|
||||
label[idx + 1] = 0;
|
||||
|
||||
if (row + 1 < H) {
|
||||
if (img[idx + W + 1])
|
||||
label[idx + W + 1] = y;
|
||||
else
|
||||
label[idx + W + 1] = 0;
|
||||
}
|
||||
}
|
||||
|
||||
if (row + 1 < H) {
|
||||
if (img[idx + W])
|
||||
label[idx + W] = y;
|
||||
else
|
||||
label[idx + W] = 0;
|
||||
}
|
||||
}
|
||||
|
||||
__global__ void init_counting(
|
||||
const int32_t* label,
|
||||
int32_t* count_init,
|
||||
const int32_t W,
|
||||
const int32_t H) {
|
||||
const uint32_t row = (blockIdx.y * blockDim.y + threadIdx.y);
|
||||
const uint32_t col = (blockIdx.x * blockDim.x + threadIdx.x);
|
||||
const uint32_t idx = row * W + col;
|
||||
|
||||
if (row >= H || col >= W)
|
||||
return;
|
||||
|
||||
int32_t y = label[idx];
|
||||
if (y > 0) {
|
||||
int32_t count_idx = y - 1;
|
||||
atomicAdd(count_init + count_idx, 1);
|
||||
}
|
||||
}
|
||||
|
||||
__global__ void final_counting(
|
||||
const int32_t* label,
|
||||
const int32_t* count_init,
|
||||
int32_t* count_final,
|
||||
const int32_t W,
|
||||
const int32_t H) {
|
||||
const uint32_t row = (blockIdx.y * blockDim.y + threadIdx.y);
|
||||
const uint32_t col = (blockIdx.x * blockDim.x + threadIdx.x);
|
||||
const uint32_t idx = row * W + col;
|
||||
|
||||
if (row >= H || col >= W)
|
||||
return;
|
||||
|
||||
int32_t y = label[idx];
|
||||
if (y > 0) {
|
||||
int32_t count_idx = y - 1;
|
||||
count_final[idx] = count_init[count_idx];
|
||||
} else {
|
||||
count_final[idx] = 0;
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace cc2d
|
||||
|
||||
std::vector<torch::Tensor> get_connected_componnets(
|
||||
const torch::Tensor& inputs) {
|
||||
AT_ASSERTM(inputs.is_cuda(), "inputs must be a CUDA tensor");
|
||||
AT_ASSERTM(inputs.ndimension() == 4, "inputs must be [N, 1, H, W] shape");
|
||||
AT_ASSERTM(
|
||||
inputs.scalar_type() == torch::kUInt8, "inputs must be a uint8 type");
|
||||
|
||||
const uint32_t N = inputs.size(0);
|
||||
const uint32_t C = inputs.size(1);
|
||||
const uint32_t H = inputs.size(2);
|
||||
const uint32_t W = inputs.size(3);
|
||||
|
||||
AT_ASSERTM(C == 1, "inputs must be [N, 1, H, W] shape");
|
||||
AT_ASSERTM((H % 2) == 0, "height must be a even number");
|
||||
AT_ASSERTM((W % 2) == 0, "width must be a even number");
|
||||
|
||||
// label must be uint32_t
|
||||
auto label_options =
|
||||
torch::TensorOptions().dtype(torch::kInt32).device(inputs.device());
|
||||
torch::Tensor labels = torch::zeros({N, C, H, W}, label_options);
|
||||
torch::Tensor counts_init = torch::zeros({N, C, H, W}, label_options);
|
||||
torch::Tensor counts_final = torch::zeros({N, C, H, W}, label_options);
|
||||
|
||||
dim3 grid = dim3(
|
||||
((W + 1) / 2 + BLOCK_COLS - 1) / BLOCK_COLS,
|
||||
((H + 1) / 2 + BLOCK_ROWS - 1) / BLOCK_ROWS);
|
||||
dim3 block = dim3(BLOCK_COLS, BLOCK_ROWS);
|
||||
dim3 grid_count =
|
||||
dim3((W + BLOCK_COLS) / BLOCK_COLS, (H + BLOCK_ROWS) / BLOCK_ROWS);
|
||||
dim3 block_count = dim3(BLOCK_COLS, BLOCK_ROWS);
|
||||
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
|
||||
|
||||
for (int n = 0; n < N; n++) {
|
||||
uint32_t offset = n * H * W;
|
||||
|
||||
cc2d::init_labeling<<<grid, block, 0, stream>>>(
|
||||
labels.data_ptr<int32_t>() + offset, W, H);
|
||||
cc2d::merge<<<grid, block, 0, stream>>>(
|
||||
inputs.data_ptr<uint8_t>() + offset,
|
||||
labels.data_ptr<int32_t>() + offset,
|
||||
W,
|
||||
H);
|
||||
cc2d::compression<<<grid, block, 0, stream>>>(
|
||||
labels.data_ptr<int32_t>() + offset, W, H);
|
||||
cc2d::final_labeling<<<grid, block, 0, stream>>>(
|
||||
inputs.data_ptr<uint8_t>() + offset,
|
||||
labels.data_ptr<int32_t>() + offset,
|
||||
W,
|
||||
H);
|
||||
|
||||
// get the counting of each pixel
|
||||
cc2d::init_counting<<<grid_count, block_count, 0, stream>>>(
|
||||
labels.data_ptr<int32_t>() + offset,
|
||||
counts_init.data_ptr<int32_t>() + offset,
|
||||
W,
|
||||
H);
|
||||
cc2d::final_counting<<<grid_count, block_count, 0, stream>>>(
|
||||
labels.data_ptr<int32_t>() + offset,
|
||||
counts_init.data_ptr<int32_t>() + offset,
|
||||
counts_final.data_ptr<int32_t>() + offset,
|
||||
W,
|
||||
H);
|
||||
}
|
||||
|
||||
// returned values are [labels, counts]
|
||||
std::vector<torch::Tensor> outputs;
|
||||
outputs.push_back(labels);
|
||||
outputs.push_back(counts_final);
|
||||
return outputs;
|
||||
}
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.def(
|
||||
"get_connected_componnets",
|
||||
&get_connected_componnets,
|
||||
"get_connected_componnets");
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
# 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.
|
||||
@@ -0,0 +1,5 @@
|
||||
# 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.
|
||||
@@ -0,0 +1,295 @@
|
||||
# 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.
|
||||
|
||||
from functools import partial
|
||||
from typing import List, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from model.segment_anything_2.sam2.modeling.backbones.utils import (
|
||||
PatchEmbed,
|
||||
window_partition,
|
||||
window_unpartition,
|
||||
)
|
||||
|
||||
from model.segment_anything_2.sam2.modeling.sam2_utils import DropPath, MLP
|
||||
|
||||
|
||||
def do_pool(x: torch.Tensor, pool: nn.Module, norm: nn.Module = None) -> torch.Tensor:
|
||||
if pool is None:
|
||||
return x
|
||||
# (B, H, W, C) -> (B, C, H, W)
|
||||
x = x.permute(0, 3, 1, 2)
|
||||
x = pool(x)
|
||||
# (B, C, H', W') -> (B, H', W', C)
|
||||
x = x.permute(0, 2, 3, 1)
|
||||
if norm:
|
||||
x = norm(x)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class MultiScaleAttention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
dim_out: int,
|
||||
num_heads: int,
|
||||
q_pool: nn.Module = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.dim = dim
|
||||
self.dim_out = dim_out
|
||||
|
||||
self.num_heads = num_heads
|
||||
head_dim = dim_out // num_heads
|
||||
self.scale = head_dim**-0.5
|
||||
|
||||
self.q_pool = q_pool
|
||||
self.qkv = nn.Linear(dim, dim_out * 3)
|
||||
self.proj = nn.Linear(dim_out, dim_out)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
B, H, W, _ = x.shape
|
||||
# qkv with shape (B, H * W, 3, nHead, C)
|
||||
qkv = self.qkv(x).reshape(B, H * W, 3, self.num_heads, -1)
|
||||
# q, k, v with shape (B, H * W, nheads, C)
|
||||
q, k, v = torch.unbind(qkv, 2)
|
||||
|
||||
# Q pooling (for downsample at stage changes)
|
||||
if self.q_pool:
|
||||
q = do_pool(q.reshape(B, H, W, -1), self.q_pool)
|
||||
H, W = q.shape[1:3] # downsampled shape
|
||||
q = q.reshape(B, H * W, self.num_heads, -1)
|
||||
|
||||
# Torch's SDPA expects [B, nheads, H*W, C] so we transpose
|
||||
x = F.scaled_dot_product_attention(
|
||||
q.transpose(1, 2),
|
||||
k.transpose(1, 2),
|
||||
v.transpose(1, 2),
|
||||
)
|
||||
# Transpose back
|
||||
x = x.transpose(1, 2)
|
||||
x = x.reshape(B, H, W, -1)
|
||||
|
||||
x = self.proj(x)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class MultiScaleBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
dim_out: int,
|
||||
num_heads: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
drop_path: float = 0.0,
|
||||
norm_layer: Union[nn.Module, str] = "LayerNorm",
|
||||
q_stride: Tuple[int, int] = None,
|
||||
act_layer: nn.Module = nn.GELU,
|
||||
window_size: int = 0,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
if isinstance(norm_layer, str):
|
||||
norm_layer = partial(getattr(nn, norm_layer), eps=1e-6)
|
||||
|
||||
self.dim = dim
|
||||
self.dim_out = dim_out
|
||||
self.norm1 = norm_layer(dim)
|
||||
|
||||
self.window_size = window_size
|
||||
|
||||
self.pool, self.q_stride = None, q_stride
|
||||
if self.q_stride:
|
||||
self.pool = nn.MaxPool2d(
|
||||
kernel_size=q_stride, stride=q_stride, ceil_mode=False
|
||||
)
|
||||
|
||||
self.attn = MultiScaleAttention(
|
||||
dim,
|
||||
dim_out,
|
||||
num_heads=num_heads,
|
||||
q_pool=self.pool,
|
||||
)
|
||||
self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()
|
||||
|
||||
self.norm2 = norm_layer(dim_out)
|
||||
self.mlp = MLP(
|
||||
dim_out,
|
||||
int(dim_out * mlp_ratio),
|
||||
dim_out,
|
||||
num_layers=2,
|
||||
activation=act_layer,
|
||||
)
|
||||
|
||||
if dim != dim_out:
|
||||
self.proj = nn.Linear(dim, dim_out)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
shortcut = x # B, H, W, C
|
||||
x = self.norm1(x)
|
||||
|
||||
# Skip connection
|
||||
if self.dim != self.dim_out:
|
||||
shortcut = do_pool(self.proj(x), self.pool)
|
||||
|
||||
# Window partition
|
||||
window_size = self.window_size
|
||||
if window_size > 0:
|
||||
H, W = x.shape[1], x.shape[2]
|
||||
x, pad_hw = window_partition(x, window_size)
|
||||
|
||||
# Window Attention + Q Pooling (if stage change)
|
||||
x = self.attn(x)
|
||||
if self.q_stride:
|
||||
# Shapes have changed due to Q pooling
|
||||
window_size = self.window_size // self.q_stride[0]
|
||||
H, W = shortcut.shape[1:3]
|
||||
|
||||
pad_h = (window_size - H % window_size) % window_size
|
||||
pad_w = (window_size - W % window_size) % window_size
|
||||
pad_hw = (H + pad_h, W + pad_w)
|
||||
|
||||
# Reverse window partition
|
||||
if self.window_size > 0:
|
||||
x = window_unpartition(x, window_size, pad_hw, (H, W))
|
||||
|
||||
x = shortcut + self.drop_path(x)
|
||||
# MLP
|
||||
x = x + self.drop_path(self.mlp(self.norm2(x)))
|
||||
return x
|
||||
|
||||
|
||||
class Hiera(nn.Module):
|
||||
"""
|
||||
Reference: https://arxiv.org/abs/2306.00989
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
embed_dim: int = 96, # initial embed dim
|
||||
num_heads: int = 1, # initial number of heads
|
||||
drop_path_rate: float = 0.0, # stochastic depth
|
||||
q_pool: int = 3, # number of q_pool stages
|
||||
q_stride: Tuple[int, int] = (2, 2), # downsample stride bet. stages
|
||||
stages: Tuple[int, ...] = (2, 3, 16, 3), # blocks per stage
|
||||
dim_mul: float = 2.0, # dim_mul factor at stage shift
|
||||
head_mul: float = 2.0, # head_mul factor at stage shift
|
||||
window_pos_embed_bkg_spatial_size: Tuple[int, int] = (14, 14),
|
||||
# window size per stage, when not using global att.
|
||||
window_spec: Tuple[int, ...] = (
|
||||
8,
|
||||
4,
|
||||
14,
|
||||
7,
|
||||
),
|
||||
# global attn in these blocks
|
||||
global_att_blocks: Tuple[int, ...] = (
|
||||
12,
|
||||
16,
|
||||
20,
|
||||
),
|
||||
return_interm_layers=True, # return feats from every stage
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
assert len(stages) == len(window_spec)
|
||||
self.window_spec = window_spec
|
||||
|
||||
depth = sum(stages)
|
||||
self.q_stride = q_stride
|
||||
self.stage_ends = [sum(stages[:i]) - 1 for i in range(1, len(stages) + 1)]
|
||||
assert 0 <= q_pool <= len(self.stage_ends[:-1])
|
||||
self.q_pool_blocks = [x + 1 for x in self.stage_ends[:-1]][:q_pool]
|
||||
self.return_interm_layers = return_interm_layers
|
||||
|
||||
self.patch_embed = PatchEmbed(
|
||||
embed_dim=embed_dim,
|
||||
)
|
||||
# Which blocks have global att?
|
||||
self.global_att_blocks = global_att_blocks
|
||||
|
||||
# Windowed positional embedding (https://arxiv.org/abs/2311.05613)
|
||||
self.window_pos_embed_bkg_spatial_size = window_pos_embed_bkg_spatial_size
|
||||
self.pos_embed = nn.Parameter(
|
||||
torch.zeros(1, embed_dim, *self.window_pos_embed_bkg_spatial_size)
|
||||
)
|
||||
self.pos_embed_window = nn.Parameter(
|
||||
torch.zeros(1, embed_dim, self.window_spec[0], self.window_spec[0])
|
||||
)
|
||||
|
||||
dpr = [
|
||||
x.item() for x in torch.linspace(0, drop_path_rate, depth)
|
||||
] # stochastic depth decay rule
|
||||
|
||||
cur_stage = 1
|
||||
self.blocks = nn.ModuleList()
|
||||
|
||||
for i in range(depth):
|
||||
dim_out = embed_dim
|
||||
# lags by a block, so first block of
|
||||
# next stage uses an initial window size
|
||||
# of previous stage and final window size of current stage
|
||||
window_size = self.window_spec[cur_stage - 1]
|
||||
|
||||
if self.global_att_blocks is not None:
|
||||
window_size = 0 if i in self.global_att_blocks else window_size
|
||||
|
||||
if i - 1 in self.stage_ends:
|
||||
dim_out = int(embed_dim * dim_mul)
|
||||
num_heads = int(num_heads * head_mul)
|
||||
cur_stage += 1
|
||||
|
||||
block = MultiScaleBlock(
|
||||
dim=embed_dim,
|
||||
dim_out=dim_out,
|
||||
num_heads=num_heads,
|
||||
drop_path=dpr[i],
|
||||
q_stride=self.q_stride if i in self.q_pool_blocks else None,
|
||||
window_size=window_size,
|
||||
)
|
||||
|
||||
embed_dim = dim_out
|
||||
self.blocks.append(block)
|
||||
|
||||
self.channel_list = (
|
||||
[self.blocks[i].dim_out for i in self.stage_ends[::-1]]
|
||||
if return_interm_layers
|
||||
else [self.blocks[-1].dim_out]
|
||||
)
|
||||
|
||||
def _get_pos_embed(self, hw: Tuple[int, int]) -> torch.Tensor:
|
||||
h, w = hw
|
||||
window_embed = self.pos_embed_window
|
||||
pos_embed = F.interpolate(self.pos_embed, size=(h, w), mode="bicubic")
|
||||
pos_embed = pos_embed + window_embed.tile(
|
||||
[x // y for x, y in zip(pos_embed.shape, window_embed.shape)]
|
||||
)
|
||||
pos_embed = pos_embed.permute(0, 2, 3, 1)
|
||||
return pos_embed
|
||||
|
||||
def forward(self, x: torch.Tensor) -> List[torch.Tensor]:
|
||||
x = self.patch_embed(x)
|
||||
# x: (B, H, W, C)
|
||||
|
||||
# Add pos embed
|
||||
x = x + self._get_pos_embed(x.shape[1:3])
|
||||
|
||||
outputs = []
|
||||
for i, blk in enumerate(self.blocks):
|
||||
x = blk(x)
|
||||
if (i == self.stage_ends[-1]) or (
|
||||
i in self.stage_ends and self.return_interm_layers
|
||||
):
|
||||
feats = x.permute(0, 3, 1, 2)
|
||||
outputs.append(feats)
|
||||
|
||||
return outputs
|
||||
@@ -0,0 +1,133 @@
|
||||
# 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.
|
||||
|
||||
from typing import List, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class ImageEncoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
trunk: nn.Module,
|
||||
neck: nn.Module,
|
||||
scalp: int = 0,
|
||||
):
|
||||
super().__init__()
|
||||
self.trunk = trunk
|
||||
self.neck = neck
|
||||
self.scalp = scalp
|
||||
assert (
|
||||
self.trunk.channel_list == self.neck.backbone_channel_list
|
||||
), f"Channel dims of trunk and neck do not match. Trunk: {self.trunk.channel_list}, neck: {self.neck.backbone_channel_list}"
|
||||
|
||||
def forward(self, sample: torch.Tensor):
|
||||
# Forward through backbone
|
||||
features, pos = self.neck(self.trunk(sample))
|
||||
if self.scalp > 0:
|
||||
# Discard the lowest resolution features
|
||||
features, pos = features[: -self.scalp], pos[: -self.scalp]
|
||||
|
||||
src = features[-1]
|
||||
output = {
|
||||
"vision_features": src,
|
||||
"vision_pos_enc": pos,
|
||||
"backbone_fpn": features,
|
||||
}
|
||||
return output
|
||||
|
||||
|
||||
class FpnNeck(nn.Module):
|
||||
"""
|
||||
A modified variant of Feature Pyramid Network (FPN) neck
|
||||
(we remove output conv and also do bicubic interpolation similar to ViT
|
||||
pos embed interpolation)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
position_encoding: nn.Module,
|
||||
d_model: int,
|
||||
backbone_channel_list: List[int],
|
||||
kernel_size: int = 1,
|
||||
stride: int = 1,
|
||||
padding: int = 0,
|
||||
fpn_interp_model: str = "bilinear",
|
||||
fuse_type: str = "sum",
|
||||
fpn_top_down_levels: Optional[List[int]] = None,
|
||||
):
|
||||
"""Initialize the neck
|
||||
:param trunk: the backbone
|
||||
:param position_encoding: the positional encoding to use
|
||||
:param d_model: the dimension of the model
|
||||
:param neck_norm: the normalization to use
|
||||
"""
|
||||
super().__init__()
|
||||
self.position_encoding = position_encoding
|
||||
self.convs = nn.ModuleList()
|
||||
self.backbone_channel_list = backbone_channel_list
|
||||
for dim in backbone_channel_list:
|
||||
current = nn.Sequential()
|
||||
current.add_module(
|
||||
"conv",
|
||||
nn.Conv2d(
|
||||
in_channels=dim,
|
||||
out_channels=d_model,
|
||||
kernel_size=kernel_size,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
),
|
||||
)
|
||||
|
||||
self.convs.append(current)
|
||||
self.fpn_interp_model = fpn_interp_model
|
||||
assert fuse_type in ["sum", "avg"]
|
||||
self.fuse_type = fuse_type
|
||||
|
||||
# levels to have top-down features in its outputs
|
||||
# e.g. if fpn_top_down_levels is [2, 3], then only outputs of level 2 and 3
|
||||
# have top-down propagation, while outputs of level 0 and level 1 have only
|
||||
# lateral features from the same backbone level.
|
||||
if fpn_top_down_levels is None:
|
||||
# default is to have top-down features on all levels
|
||||
fpn_top_down_levels = range(len(self.convs))
|
||||
self.fpn_top_down_levels = list(fpn_top_down_levels)
|
||||
|
||||
def forward(self, xs: List[torch.Tensor]):
|
||||
|
||||
out = [None] * len(self.convs)
|
||||
pos = [None] * len(self.convs)
|
||||
assert len(xs) == len(self.convs)
|
||||
# fpn forward pass
|
||||
# see https://github.com/facebookresearch/detectron2/blob/main/detectron2/modeling/backbone/fpn.py
|
||||
prev_features = None
|
||||
# forward in top-down order (from low to high resolution)
|
||||
n = len(self.convs) - 1
|
||||
for i in range(n, -1, -1):
|
||||
x = xs[i]
|
||||
lateral_features = self.convs[n - i](x)
|
||||
if i in self.fpn_top_down_levels and prev_features is not None:
|
||||
top_down_features = F.interpolate(
|
||||
prev_features.to(dtype=torch.float32),
|
||||
scale_factor=2.0,
|
||||
mode=self.fpn_interp_model,
|
||||
align_corners=(
|
||||
None if self.fpn_interp_model == "nearest" else False
|
||||
),
|
||||
antialias=False,
|
||||
)
|
||||
prev_features = lateral_features + top_down_features
|
||||
if self.fuse_type == "avg":
|
||||
prev_features /= 2
|
||||
else:
|
||||
prev_features = lateral_features
|
||||
x_out = prev_features
|
||||
out[i] = x_out
|
||||
pos[i] = self.position_encoding(x_out).to(x_out.dtype)
|
||||
|
||||
return out, pos
|
||||
@@ -0,0 +1,95 @@
|
||||
# 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.
|
||||
|
||||
"""Some utilities for backbones, in particular for windowing"""
|
||||
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
def window_partition(x, window_size):
|
||||
"""
|
||||
Partition into non-overlapping windows with padding if needed.
|
||||
Args:
|
||||
x (tensor): input tokens with [B, H, W, C].
|
||||
window_size (int): window size.
|
||||
Returns:
|
||||
windows: windows after partition with [B * num_windows, window_size, window_size, C].
|
||||
(Hp, Wp): padded height and width before partition
|
||||
"""
|
||||
B, H, W, C = x.shape
|
||||
|
||||
pad_h = (window_size - H % window_size) % window_size
|
||||
pad_w = (window_size - W % window_size) % window_size
|
||||
if pad_h > 0 or pad_w > 0:
|
||||
x = F.pad(x, (0, 0, 0, pad_w, 0, pad_h))
|
||||
Hp, Wp = H + pad_h, W + pad_w
|
||||
|
||||
x = x.view(B, Hp // window_size, window_size, Wp // window_size, window_size, C)
|
||||
windows = (
|
||||
x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C)
|
||||
)
|
||||
return windows, (Hp, Wp)
|
||||
|
||||
|
||||
def window_unpartition(windows, window_size, pad_hw, hw):
|
||||
"""
|
||||
Window unpartition into original sequences and removing padding.
|
||||
Args:
|
||||
x (tensor): input tokens with [B * num_windows, window_size, window_size, C].
|
||||
window_size (int): window size.
|
||||
pad_hw (Tuple): padded height and width (Hp, Wp).
|
||||
hw (Tuple): original height and width (H, W) before padding.
|
||||
Returns:
|
||||
x: unpartitioned sequences with [B, H, W, C].
|
||||
"""
|
||||
Hp, Wp = pad_hw
|
||||
H, W = hw
|
||||
B = windows.shape[0] // (Hp * Wp // window_size // window_size)
|
||||
x = windows.view(
|
||||
B, Hp // window_size, Wp // window_size, window_size, window_size, -1
|
||||
)
|
||||
x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, Hp, Wp, -1)
|
||||
|
||||
if Hp > H or Wp > W:
|
||||
x = x[:, :H, :W, :].contiguous()
|
||||
return x
|
||||
|
||||
|
||||
class PatchEmbed(nn.Module):
|
||||
"""
|
||||
Image to Patch Embedding.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
kernel_size: Tuple[int, ...] = (7, 7),
|
||||
stride: Tuple[int, ...] = (4, 4),
|
||||
padding: Tuple[int, ...] = (3, 3),
|
||||
in_chans: int = 3,
|
||||
embed_dim: int = 768,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
kernel_size (Tuple): kernel size of the projection layer.
|
||||
stride (Tuple): stride of the projection layer.
|
||||
padding (Tuple): padding size of the projection layer.
|
||||
in_chans (int): Number of input image channels.
|
||||
embed_dim (int): embed_dim (int): Patch embedding dimension.
|
||||
"""
|
||||
super().__init__()
|
||||
self.proj = nn.Conv2d(
|
||||
in_chans, embed_dim, kernel_size=kernel_size, stride=stride, padding=padding
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = self.proj(x)
|
||||
# B C H W -> B H W C
|
||||
x = x.permute(0, 2, 3, 1)
|
||||
return x
|
||||
@@ -0,0 +1,169 @@
|
||||
# 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.
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
from torch import nn, Tensor
|
||||
|
||||
from model.segment_anything_2.sam2.modeling.sam.transformer import RoPEAttention
|
||||
|
||||
from model.segment_anything_2.sam2.modeling.sam2_utils import get_activation_fn, get_clones
|
||||
|
||||
|
||||
class MemoryAttentionLayer(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
activation: str,
|
||||
cross_attention: nn.Module,
|
||||
d_model: int,
|
||||
dim_feedforward: int,
|
||||
dropout: float,
|
||||
pos_enc_at_attn: bool,
|
||||
pos_enc_at_cross_attn_keys: bool,
|
||||
pos_enc_at_cross_attn_queries: bool,
|
||||
self_attention: nn.Module,
|
||||
):
|
||||
super().__init__()
|
||||
self.d_model = d_model
|
||||
self.dim_feedforward = dim_feedforward
|
||||
self.dropout_value = dropout
|
||||
self.self_attn = self_attention
|
||||
self.cross_attn_image = cross_attention
|
||||
|
||||
# Implementation of Feedforward model
|
||||
self.linear1 = nn.Linear(d_model, dim_feedforward)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.linear2 = nn.Linear(dim_feedforward, d_model)
|
||||
|
||||
self.norm1 = nn.LayerNorm(d_model)
|
||||
self.norm2 = nn.LayerNorm(d_model)
|
||||
self.norm3 = nn.LayerNorm(d_model)
|
||||
self.dropout1 = nn.Dropout(dropout)
|
||||
self.dropout2 = nn.Dropout(dropout)
|
||||
self.dropout3 = nn.Dropout(dropout)
|
||||
|
||||
self.activation_str = activation
|
||||
self.activation = get_activation_fn(activation)
|
||||
|
||||
# Where to add pos enc
|
||||
self.pos_enc_at_attn = pos_enc_at_attn
|
||||
self.pos_enc_at_cross_attn_queries = pos_enc_at_cross_attn_queries
|
||||
self.pos_enc_at_cross_attn_keys = pos_enc_at_cross_attn_keys
|
||||
|
||||
def _forward_sa(self, tgt, query_pos):
|
||||
# Self-Attention
|
||||
tgt2 = self.norm1(tgt)
|
||||
q = k = tgt2 + query_pos if self.pos_enc_at_attn else tgt2
|
||||
tgt2 = self.self_attn(q, k, v=tgt2)
|
||||
tgt = tgt + self.dropout1(tgt2)
|
||||
return tgt
|
||||
|
||||
def _forward_ca(self, tgt, memory, query_pos, pos, num_k_exclude_rope=0):
|
||||
kwds = {}
|
||||
if num_k_exclude_rope > 0:
|
||||
assert isinstance(self.cross_attn_image, RoPEAttention)
|
||||
kwds = {"num_k_exclude_rope": num_k_exclude_rope}
|
||||
|
||||
# Cross-Attention
|
||||
tgt2 = self.norm2(tgt)
|
||||
tgt2 = self.cross_attn_image(
|
||||
q=tgt2 + query_pos if self.pos_enc_at_cross_attn_queries else tgt2,
|
||||
k=memory + pos if self.pos_enc_at_cross_attn_keys else memory,
|
||||
v=memory,
|
||||
**kwds,
|
||||
)
|
||||
tgt = tgt + self.dropout2(tgt2)
|
||||
return tgt
|
||||
|
||||
def forward(
|
||||
self,
|
||||
tgt,
|
||||
memory,
|
||||
pos: Optional[Tensor] = None,
|
||||
query_pos: Optional[Tensor] = None,
|
||||
num_k_exclude_rope: int = 0,
|
||||
) -> torch.Tensor:
|
||||
|
||||
# Self-Attn, Cross-Attn
|
||||
tgt = self._forward_sa(tgt, query_pos)
|
||||
tgt = self._forward_ca(tgt, memory, query_pos, pos, num_k_exclude_rope)
|
||||
# MLP
|
||||
tgt2 = self.norm3(tgt)
|
||||
tgt2 = self.linear2(self.dropout(self.activation(self.linear1(tgt2))))
|
||||
tgt = tgt + self.dropout3(tgt2)
|
||||
return tgt
|
||||
|
||||
|
||||
class MemoryAttention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
d_model: int,
|
||||
pos_enc_at_input: bool,
|
||||
layer: nn.Module,
|
||||
num_layers: int,
|
||||
batch_first: bool = True, # Do layers expect batch first input?
|
||||
):
|
||||
super().__init__()
|
||||
self.d_model = d_model
|
||||
self.layers = get_clones(layer, num_layers)
|
||||
self.num_layers = num_layers
|
||||
self.norm = nn.LayerNorm(d_model)
|
||||
self.pos_enc_at_input = pos_enc_at_input
|
||||
self.batch_first = batch_first
|
||||
|
||||
def forward(
|
||||
self,
|
||||
curr: torch.Tensor, # self-attention inputs
|
||||
memory: torch.Tensor, # cross-attention inputs
|
||||
curr_pos: Optional[Tensor] = None, # pos_enc for self-attention inputs
|
||||
memory_pos: Optional[Tensor] = None, # pos_enc for cross-attention inputs
|
||||
num_obj_ptr_tokens: int = 0, # number of object pointer *tokens*
|
||||
):
|
||||
if isinstance(curr, list):
|
||||
assert isinstance(curr_pos, list)
|
||||
assert len(curr) == len(curr_pos) == 1
|
||||
curr, curr_pos = (
|
||||
curr[0],
|
||||
curr_pos[0],
|
||||
)
|
||||
|
||||
assert (
|
||||
curr.shape[1] == memory.shape[1]
|
||||
), "Batch size must be the same for curr and memory"
|
||||
|
||||
output = curr
|
||||
if self.pos_enc_at_input and curr_pos is not None:
|
||||
output = output + 0.1 * curr_pos
|
||||
|
||||
if self.batch_first:
|
||||
# Convert to batch first
|
||||
output = output.transpose(0, 1)
|
||||
curr_pos = curr_pos.transpose(0, 1)
|
||||
memory = memory.transpose(0, 1)
|
||||
memory_pos = memory_pos.transpose(0, 1)
|
||||
|
||||
for layer in self.layers:
|
||||
kwds = {}
|
||||
if isinstance(layer.cross_attn_image, RoPEAttention):
|
||||
kwds = {"num_k_exclude_rope": num_obj_ptr_tokens}
|
||||
|
||||
output = layer(
|
||||
tgt=output,
|
||||
memory=memory,
|
||||
pos=memory_pos,
|
||||
query_pos=curr_pos,
|
||||
**kwds,
|
||||
)
|
||||
normed_output = self.norm(output)
|
||||
|
||||
if self.batch_first:
|
||||
# Convert back to seq first
|
||||
normed_output = normed_output.transpose(0, 1)
|
||||
curr_pos = curr_pos.transpose(0, 1)
|
||||
|
||||
return normed_output
|
||||
@@ -0,0 +1,182 @@
|
||||
# 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 math
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from model.segment_anything_2.sam2.modeling.sam2_utils import DropPath, get_clones, LayerNorm2d
|
||||
|
||||
|
||||
class MaskDownSampler(nn.Module):
|
||||
"""
|
||||
Progressively downsample a mask by total_stride, each time by stride.
|
||||
Note that LayerNorm is applied per *token*, like in ViT.
|
||||
|
||||
With each downsample (by a factor stride**2), channel capacity increases by the same factor.
|
||||
In the end, we linearly project to embed_dim channels.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
embed_dim=256,
|
||||
kernel_size=4,
|
||||
stride=4,
|
||||
padding=0,
|
||||
total_stride=16,
|
||||
activation=nn.GELU,
|
||||
):
|
||||
super().__init__()
|
||||
num_layers = int(math.log2(total_stride) // math.log2(stride))
|
||||
assert stride**num_layers == total_stride
|
||||
self.encoder = nn.Sequential()
|
||||
mask_in_chans, mask_out_chans = 1, 1
|
||||
for _ in range(num_layers):
|
||||
mask_out_chans = mask_in_chans * (stride**2)
|
||||
self.encoder.append(
|
||||
nn.Conv2d(
|
||||
mask_in_chans,
|
||||
mask_out_chans,
|
||||
kernel_size=kernel_size,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
)
|
||||
)
|
||||
self.encoder.append(LayerNorm2d(mask_out_chans))
|
||||
self.encoder.append(activation())
|
||||
mask_in_chans = mask_out_chans
|
||||
|
||||
self.encoder.append(nn.Conv2d(mask_out_chans, embed_dim, kernel_size=1))
|
||||
|
||||
def forward(self, x):
|
||||
return self.encoder(x)
|
||||
|
||||
|
||||
# Lightly adapted from ConvNext (https://github.com/facebookresearch/ConvNeXt)
|
||||
class CXBlock(nn.Module):
|
||||
r"""ConvNeXt Block. There are two equivalent implementations:
|
||||
(1) DwConv -> LayerNorm (channels_first) -> 1x1 Conv -> GELU -> 1x1 Conv; all in (N, C, H, W)
|
||||
(2) DwConv -> Permute to (N, H, W, C); LayerNorm (channels_last) -> Linear -> GELU -> Linear; Permute back
|
||||
We use (2) as we find it slightly faster in PyTorch
|
||||
|
||||
Args:
|
||||
dim (int): Number of input channels.
|
||||
drop_path (float): Stochastic depth rate. Default: 0.0
|
||||
layer_scale_init_value (float): Init value for Layer Scale. Default: 1e-6.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
kernel_size=7,
|
||||
padding=3,
|
||||
drop_path=0.0,
|
||||
layer_scale_init_value=1e-6,
|
||||
use_dwconv=True,
|
||||
):
|
||||
super().__init__()
|
||||
self.dwconv = nn.Conv2d(
|
||||
dim,
|
||||
dim,
|
||||
kernel_size=kernel_size,
|
||||
padding=padding,
|
||||
groups=dim if use_dwconv else 1,
|
||||
) # depthwise conv
|
||||
self.norm = LayerNorm2d(dim, eps=1e-6)
|
||||
self.pwconv1 = nn.Linear(
|
||||
dim, 4 * dim
|
||||
) # pointwise/1x1 convs, implemented with linear layers
|
||||
self.act = nn.GELU()
|
||||
self.pwconv2 = nn.Linear(4 * dim, dim)
|
||||
# modified by ZhangYx from self.gamma to self.weight. Due to (https://github.com/facebookresearch/segment-anything-2/issues/85)
|
||||
self.weight = (
|
||||
nn.Parameter(layer_scale_init_value * torch.ones((dim)), requires_grad=True)
|
||||
if layer_scale_init_value > 0
|
||||
else None
|
||||
)
|
||||
self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()
|
||||
|
||||
def forward(self, x):
|
||||
input = x
|
||||
x = self.dwconv(x)
|
||||
x = self.norm(x)
|
||||
x = x.permute(0, 2, 3, 1) # (N, C, H, W) -> (N, H, W, C)
|
||||
x = self.pwconv1(x)
|
||||
x = self.act(x)
|
||||
x = self.pwconv2(x)
|
||||
if self.weight is not None:
|
||||
x = self.weight * x
|
||||
x = x.permute(0, 3, 1, 2) # (N, H, W, C) -> (N, C, H, W)
|
||||
|
||||
x = input + self.drop_path(x)
|
||||
return x
|
||||
|
||||
|
||||
class Fuser(nn.Module):
|
||||
def __init__(self, layer, num_layers, dim=None, input_projection=False):
|
||||
super().__init__()
|
||||
self.proj = nn.Identity()
|
||||
self.layers = get_clones(layer, num_layers)
|
||||
|
||||
if input_projection:
|
||||
assert dim is not None
|
||||
self.proj = nn.Conv2d(dim, dim, kernel_size=1)
|
||||
|
||||
def forward(self, x):
|
||||
# normally x: (N, C, H, W)
|
||||
x = self.proj(x)
|
||||
for layer in self.layers:
|
||||
x = layer(x)
|
||||
return x
|
||||
|
||||
|
||||
class MemoryEncoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
out_dim,
|
||||
mask_downsampler,
|
||||
fuser,
|
||||
position_encoding,
|
||||
in_dim=256, # in_dim of pix_feats
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.mask_downsampler = mask_downsampler
|
||||
|
||||
self.pix_feat_proj = nn.Conv2d(in_dim, in_dim, kernel_size=1)
|
||||
self.fuser = fuser
|
||||
self.position_encoding = position_encoding
|
||||
self.out_proj = nn.Identity()
|
||||
if out_dim != in_dim:
|
||||
self.out_proj = nn.Conv2d(in_dim, out_dim, kernel_size=1)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
pix_feat: torch.Tensor,
|
||||
masks: torch.Tensor,
|
||||
skip_mask_sigmoid: bool = False,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
## Process masks
|
||||
# sigmoid, so that less domain shift from gt masks which are bool
|
||||
if not skip_mask_sigmoid:
|
||||
masks = F.sigmoid(masks)
|
||||
masks = self.mask_downsampler(masks)
|
||||
|
||||
## Fuse pix_feats and downsampled masks
|
||||
# in case the visual features are on CPU, cast them to CUDA
|
||||
pix_feat = pix_feat.to(masks.device)
|
||||
|
||||
x = self.pix_feat_proj(pix_feat)
|
||||
x = x + masks
|
||||
x = self.fuser(x)
|
||||
x = self.out_proj(x)
|
||||
|
||||
pos = self.position_encoding(x).to(x.dtype)
|
||||
|
||||
return {"vision_features": x, "vision_pos_enc": [pos]}
|
||||
@@ -0,0 +1,216 @@
|
||||
# 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 math
|
||||
from typing import Any, Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
|
||||
class PositionEmbeddingSine(nn.Module):
|
||||
"""
|
||||
This is a more standard version of the position embedding, very similar to the one
|
||||
used by the Attention is all you need paper, generalized to work on images.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_pos_feats,
|
||||
temperature: int = 10000,
|
||||
normalize: bool = True,
|
||||
scale: Optional[float] = None,
|
||||
):
|
||||
super().__init__()
|
||||
assert num_pos_feats % 2 == 0, "Expecting even model width"
|
||||
self.num_pos_feats = num_pos_feats // 2
|
||||
self.temperature = temperature
|
||||
self.normalize = normalize
|
||||
if scale is not None and normalize is False:
|
||||
raise ValueError("normalize should be True if scale is passed")
|
||||
if scale is None:
|
||||
scale = 2 * math.pi
|
||||
self.scale = scale
|
||||
|
||||
self.cache = {}
|
||||
|
||||
def _encode_xy(self, x, y):
|
||||
# The positions are expected to be normalized
|
||||
assert len(x) == len(y) and x.ndim == y.ndim == 1
|
||||
x_embed = x * self.scale
|
||||
y_embed = y * self.scale
|
||||
|
||||
dim_t = torch.arange(self.num_pos_feats, dtype=torch.float32, device=x.device)
|
||||
dim_t = self.temperature ** (2 * (dim_t // 2) / self.num_pos_feats)
|
||||
|
||||
pos_x = x_embed[:, None] / dim_t
|
||||
pos_y = y_embed[:, None] / dim_t
|
||||
pos_x = torch.stack(
|
||||
(pos_x[:, 0::2].sin(), pos_x[:, 1::2].cos()), dim=2
|
||||
).flatten(1)
|
||||
pos_y = torch.stack(
|
||||
(pos_y[:, 0::2].sin(), pos_y[:, 1::2].cos()), dim=2
|
||||
).flatten(1)
|
||||
return pos_x, pos_y
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_boxes(self, x, y, w, h):
|
||||
pos_x, pos_y = self._encode_xy(x, y)
|
||||
pos = torch.cat((pos_y, pos_x, h[:, None], w[:, None]), dim=1)
|
||||
return pos
|
||||
|
||||
encode = encode_boxes # Backwards compatibility
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_points(self, x, y, labels):
|
||||
(bx, nx), (by, ny), (bl, nl) = x.shape, y.shape, labels.shape
|
||||
assert bx == by and nx == ny and bx == bl and nx == nl
|
||||
pos_x, pos_y = self._encode_xy(x.flatten(), y.flatten())
|
||||
pos_x, pos_y = pos_x.reshape(bx, nx, -1), pos_y.reshape(by, ny, -1)
|
||||
pos = torch.cat((pos_y, pos_x, labels[:, :, None]), dim=2)
|
||||
return pos
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, x: torch.Tensor):
|
||||
cache_key = (x.shape[-2], x.shape[-1])
|
||||
if cache_key in self.cache:
|
||||
return self.cache[cache_key][None].repeat(x.shape[0], 1, 1, 1)
|
||||
y_embed = (
|
||||
torch.arange(1, x.shape[-2] + 1, dtype=torch.float32, device=x.device)
|
||||
.view(1, -1, 1)
|
||||
.repeat(x.shape[0], 1, x.shape[-1])
|
||||
)
|
||||
x_embed = (
|
||||
torch.arange(1, x.shape[-1] + 1, dtype=torch.float32, device=x.device)
|
||||
.view(1, 1, -1)
|
||||
.repeat(x.shape[0], x.shape[-2], 1)
|
||||
)
|
||||
|
||||
if self.normalize:
|
||||
eps = 1e-6
|
||||
y_embed = y_embed / (y_embed[:, -1:, :] + eps) * self.scale
|
||||
x_embed = x_embed / (x_embed[:, :, -1:] + eps) * self.scale
|
||||
|
||||
dim_t = torch.arange(self.num_pos_feats, dtype=torch.float32, device=x.device)
|
||||
dim_t = self.temperature ** (2 * (dim_t // 2) / self.num_pos_feats)
|
||||
|
||||
pos_x = x_embed[:, :, :, None] / dim_t
|
||||
pos_y = y_embed[:, :, :, None] / dim_t
|
||||
pos_x = torch.stack(
|
||||
(pos_x[:, :, :, 0::2].sin(), pos_x[:, :, :, 1::2].cos()), dim=4
|
||||
).flatten(3)
|
||||
pos_y = torch.stack(
|
||||
(pos_y[:, :, :, 0::2].sin(), pos_y[:, :, :, 1::2].cos()), dim=4
|
||||
).flatten(3)
|
||||
pos = torch.cat((pos_y, pos_x), dim=3).permute(0, 3, 1, 2)
|
||||
self.cache[cache_key] = pos[0]
|
||||
return pos
|
||||
|
||||
|
||||
class PositionEmbeddingRandom(nn.Module):
|
||||
"""
|
||||
Positional encoding using random spatial frequencies.
|
||||
"""
|
||||
|
||||
def __init__(self, num_pos_feats: int = 64, scale: Optional[float] = None) -> None:
|
||||
super().__init__()
|
||||
if scale is None or scale <= 0.0:
|
||||
scale = 1.0
|
||||
self.register_buffer(
|
||||
"positional_encoding_gaussian_matrix",
|
||||
scale * torch.randn((2, num_pos_feats)),
|
||||
)
|
||||
|
||||
def _pe_encoding(self, coords: torch.Tensor) -> torch.Tensor:
|
||||
"""Positionally encode points that are normalized to [0,1]."""
|
||||
# assuming coords are in [0, 1]^2 square and have d_1 x ... x d_n x 2 shape
|
||||
coords = 2 * coords - 1
|
||||
coords = coords @ self.positional_encoding_gaussian_matrix
|
||||
coords = 2 * np.pi * coords
|
||||
# outputs d_1 x ... x d_n x C shape
|
||||
return torch.cat([torch.sin(coords), torch.cos(coords)], dim=-1)
|
||||
|
||||
def forward(self, size: Tuple[int, int]) -> torch.Tensor:
|
||||
"""Generate positional encoding for a grid of the specified size."""
|
||||
h, w = size
|
||||
device: Any = self.positional_encoding_gaussian_matrix.device
|
||||
grid = torch.ones((h, w), device=device, dtype=torch.float32)
|
||||
y_embed = grid.cumsum(dim=0) - 0.5
|
||||
x_embed = grid.cumsum(dim=1) - 0.5
|
||||
y_embed = y_embed / h
|
||||
x_embed = x_embed / w
|
||||
|
||||
pe = self._pe_encoding(torch.stack([x_embed, y_embed], dim=-1))
|
||||
return pe.permute(2, 0, 1) # C x H x W
|
||||
|
||||
def forward_with_coords(
|
||||
self, coords_input: torch.Tensor, image_size: Tuple[int, int]
|
||||
) -> torch.Tensor:
|
||||
"""Positionally encode points that are not normalized to [0,1]."""
|
||||
coords = coords_input.clone()
|
||||
coords[:, :, 0] = coords[:, :, 0] / image_size[1]
|
||||
coords[:, :, 1] = coords[:, :, 1] / image_size[0]
|
||||
return self._pe_encoding(coords.to(torch.float)) # B x N x C
|
||||
|
||||
|
||||
# Rotary Positional Encoding, adapted from:
|
||||
# 1. https://github.com/meta-llama/codellama/blob/main/llama/model.py
|
||||
# 2. https://github.com/naver-ai/rope-vit
|
||||
# 3. https://github.com/lucidrains/rotary-embedding-torch
|
||||
|
||||
|
||||
def init_t_xy(end_x: int, end_y: int):
|
||||
t = torch.arange(end_x * end_y, dtype=torch.float32)
|
||||
t_x = (t % end_x).float()
|
||||
t_y = torch.div(t, end_x, rounding_mode="floor").float()
|
||||
return t_x, t_y
|
||||
|
||||
|
||||
def compute_axial_cis(dim: int, end_x: int, end_y: int, theta: float = 10000.0):
|
||||
freqs_x = 1.0 / (theta ** (torch.arange(0, dim, 4)[: (dim // 4)].float() / dim))
|
||||
freqs_y = 1.0 / (theta ** (torch.arange(0, dim, 4)[: (dim // 4)].float() / dim))
|
||||
|
||||
t_x, t_y = init_t_xy(end_x, end_y)
|
||||
freqs_x = torch.outer(t_x, freqs_x)
|
||||
freqs_y = torch.outer(t_y, freqs_y)
|
||||
freqs_cis_x = torch.polar(torch.ones_like(freqs_x), freqs_x)
|
||||
freqs_cis_y = torch.polar(torch.ones_like(freqs_y), freqs_y)
|
||||
return torch.cat([freqs_cis_x, freqs_cis_y], dim=-1)
|
||||
|
||||
|
||||
def reshape_for_broadcast(freqs_cis: torch.Tensor, x: torch.Tensor):
|
||||
ndim = x.ndim
|
||||
assert 0 <= 1 < ndim
|
||||
assert freqs_cis.shape == (x.shape[-2], x.shape[-1])
|
||||
shape = [d if i >= ndim - 2 else 1 for i, d in enumerate(x.shape)]
|
||||
return freqs_cis.view(*shape)
|
||||
|
||||
|
||||
def apply_rotary_enc(
|
||||
xq: torch.Tensor,
|
||||
xk: torch.Tensor,
|
||||
freqs_cis: torch.Tensor,
|
||||
repeat_freqs_k: bool = False,
|
||||
):
|
||||
xq_ = torch.view_as_complex(xq.float().reshape(*xq.shape[:-1], -1, 2))
|
||||
xk_ = (
|
||||
torch.view_as_complex(xk.float().reshape(*xk.shape[:-1], -1, 2))
|
||||
if xk.shape[-2] != 0
|
||||
else None
|
||||
)
|
||||
freqs_cis = reshape_for_broadcast(freqs_cis, xq_)
|
||||
xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(3)
|
||||
if xk_ is None:
|
||||
# no keys to rotate, due to dropout
|
||||
return xq_out.type_as(xq).to(xq.device), xk
|
||||
# repeat freqs along seq_len dim to match k seq_len
|
||||
if repeat_freqs_k:
|
||||
r = xk_.shape[-2] // xq_.shape[-2]
|
||||
freqs_cis = freqs_cis.repeat(*([1] * (freqs_cis.ndim - 2)), r, 1)
|
||||
xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(3)
|
||||
return xq_out.type_as(xq).to(xq.device), xk_out.type_as(xk).to(xk.device)
|
||||
@@ -0,0 +1,5 @@
|
||||
# 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.
|
||||
@@ -0,0 +1,295 @@
|
||||
# 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.
|
||||
|
||||
from typing import List, Optional, Tuple, Type
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from model.segment_anything_2.sam2.modeling.sam2_utils import LayerNorm2d, MLP
|
||||
|
||||
|
||||
class MaskDecoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
transformer_dim: int,
|
||||
transformer: nn.Module,
|
||||
num_multimask_outputs: int = 3,
|
||||
activation: Type[nn.Module] = nn.GELU,
|
||||
iou_head_depth: int = 3,
|
||||
iou_head_hidden_dim: int = 256,
|
||||
use_high_res_features: bool = False,
|
||||
iou_prediction_use_sigmoid=False,
|
||||
dynamic_multimask_via_stability=False,
|
||||
dynamic_multimask_stability_delta=0.05,
|
||||
dynamic_multimask_stability_thresh=0.98,
|
||||
pred_obj_scores: bool = False,
|
||||
pred_obj_scores_mlp: bool = False,
|
||||
use_multimask_token_for_obj_ptr: bool = False,
|
||||
) -> None:
|
||||
"""
|
||||
Predicts masks given an image and prompt embeddings, using a
|
||||
transformer architecture.
|
||||
|
||||
Arguments:
|
||||
transformer_dim (int): the channel dimension of the transformer
|
||||
transformer (nn.Module): the transformer used to predict masks
|
||||
num_multimask_outputs (int): the number of masks to predict
|
||||
when disambiguating masks
|
||||
activation (nn.Module): the type of activation to use when
|
||||
upscaling masks
|
||||
iou_head_depth (int): the depth of the MLP used to predict
|
||||
mask quality
|
||||
iou_head_hidden_dim (int): the hidden dimension of the MLP
|
||||
used to predict mask quality
|
||||
"""
|
||||
super().__init__()
|
||||
self.transformer_dim = transformer_dim
|
||||
self.transformer = transformer
|
||||
|
||||
self.num_multimask_outputs = num_multimask_outputs
|
||||
|
||||
self.iou_token = nn.Embedding(1, transformer_dim)
|
||||
self.num_mask_tokens = num_multimask_outputs + 1
|
||||
self.mask_tokens = nn.Embedding(self.num_mask_tokens, transformer_dim)
|
||||
|
||||
self.pred_obj_scores = pred_obj_scores
|
||||
if self.pred_obj_scores:
|
||||
self.obj_score_token = nn.Embedding(1, transformer_dim)
|
||||
self.use_multimask_token_for_obj_ptr = use_multimask_token_for_obj_ptr
|
||||
|
||||
self.output_upscaling = nn.Sequential(
|
||||
nn.ConvTranspose2d(
|
||||
transformer_dim, transformer_dim // 4, kernel_size=2, stride=2
|
||||
),
|
||||
LayerNorm2d(transformer_dim // 4),
|
||||
activation(),
|
||||
nn.ConvTranspose2d(
|
||||
transformer_dim // 4, transformer_dim // 8, kernel_size=2, stride=2
|
||||
),
|
||||
activation(),
|
||||
)
|
||||
self.use_high_res_features = use_high_res_features
|
||||
if use_high_res_features:
|
||||
self.conv_s0 = nn.Conv2d(
|
||||
transformer_dim, transformer_dim // 8, kernel_size=1, stride=1
|
||||
)
|
||||
self.conv_s1 = nn.Conv2d(
|
||||
transformer_dim, transformer_dim // 4, kernel_size=1, stride=1
|
||||
)
|
||||
|
||||
self.output_hypernetworks_mlps = nn.ModuleList(
|
||||
[
|
||||
MLP(transformer_dim, transformer_dim, transformer_dim // 8, 3)
|
||||
for i in range(self.num_mask_tokens)
|
||||
]
|
||||
)
|
||||
|
||||
self.iou_prediction_head = MLP(
|
||||
transformer_dim,
|
||||
iou_head_hidden_dim,
|
||||
self.num_mask_tokens,
|
||||
iou_head_depth,
|
||||
sigmoid_output=iou_prediction_use_sigmoid,
|
||||
)
|
||||
if self.pred_obj_scores:
|
||||
self.pred_obj_score_head = nn.Linear(transformer_dim, 1)
|
||||
if pred_obj_scores_mlp:
|
||||
self.pred_obj_score_head = MLP(transformer_dim, transformer_dim, 1, 3)
|
||||
|
||||
# When outputting a single mask, optionally we can dynamically fall back to the best
|
||||
# multimask output token if the single mask output token gives low stability scores.
|
||||
self.dynamic_multimask_via_stability = dynamic_multimask_via_stability
|
||||
self.dynamic_multimask_stability_delta = dynamic_multimask_stability_delta
|
||||
self.dynamic_multimask_stability_thresh = dynamic_multimask_stability_thresh
|
||||
|
||||
def forward(
|
||||
self,
|
||||
image_embeddings: torch.Tensor,
|
||||
image_pe: torch.Tensor,
|
||||
sparse_prompt_embeddings: torch.Tensor,
|
||||
dense_prompt_embeddings: torch.Tensor,
|
||||
multimask_output: bool,
|
||||
repeat_image: bool,
|
||||
high_res_features: Optional[List[torch.Tensor]] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Predict masks given image and prompt embeddings.
|
||||
|
||||
Arguments:
|
||||
image_embeddings (torch.Tensor): the embeddings from the image encoder
|
||||
image_pe (torch.Tensor): positional encoding with the shape of image_embeddings
|
||||
sparse_prompt_embeddings (torch.Tensor): the embeddings of the points and boxes
|
||||
dense_prompt_embeddings (torch.Tensor): the embeddings of the mask inputs
|
||||
multimask_output (bool): Whether to return multiple masks or a single
|
||||
mask.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: batched predicted masks
|
||||
torch.Tensor: batched predictions of mask quality
|
||||
torch.Tensor: batched SAM token for mask output
|
||||
"""
|
||||
masks, iou_pred, mask_tokens_out, object_score_logits = self.predict_masks(
|
||||
image_embeddings=image_embeddings,
|
||||
image_pe=image_pe,
|
||||
sparse_prompt_embeddings=sparse_prompt_embeddings,
|
||||
dense_prompt_embeddings=dense_prompt_embeddings,
|
||||
repeat_image=repeat_image,
|
||||
high_res_features=high_res_features,
|
||||
)
|
||||
|
||||
# Select the correct mask or masks for output
|
||||
if multimask_output:
|
||||
masks = masks[:, 1:, :, :]
|
||||
iou_pred = iou_pred[:, 1:]
|
||||
elif self.dynamic_multimask_via_stability and not self.training:
|
||||
masks, iou_pred = self._dynamic_multimask_via_stability(masks, iou_pred)
|
||||
else:
|
||||
masks = masks[:, 0:1, :, :]
|
||||
iou_pred = iou_pred[:, 0:1]
|
||||
|
||||
if multimask_output and self.use_multimask_token_for_obj_ptr:
|
||||
sam_tokens_out = mask_tokens_out[:, 1:] # [b, 3, c] shape
|
||||
else:
|
||||
# Take the mask output token. Here we *always* use the token for single mask output.
|
||||
# At test time, even if we track after 1-click (and using multimask_output=True),
|
||||
# we still take the single mask token here. The rationale is that we always track
|
||||
# after multiple clicks during training, so the past tokens seen during training
|
||||
# are always the single mask token (and we'll let it be the object-memory token).
|
||||
sam_tokens_out = mask_tokens_out[:, 0:1] # [b, 1, c] shape
|
||||
|
||||
# Prepare output
|
||||
return masks, iou_pred, sam_tokens_out, object_score_logits
|
||||
|
||||
def predict_masks(
|
||||
self,
|
||||
image_embeddings: torch.Tensor,
|
||||
image_pe: torch.Tensor,
|
||||
sparse_prompt_embeddings: torch.Tensor,
|
||||
dense_prompt_embeddings: torch.Tensor,
|
||||
repeat_image: bool,
|
||||
high_res_features: Optional[List[torch.Tensor]] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Predicts masks. See 'forward' for more details."""
|
||||
# Concatenate output tokens
|
||||
s = 0
|
||||
if self.pred_obj_scores:
|
||||
output_tokens = torch.cat(
|
||||
[
|
||||
self.obj_score_token.weight,
|
||||
self.iou_token.weight,
|
||||
self.mask_tokens.weight,
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
s = 1
|
||||
else:
|
||||
output_tokens = torch.cat(
|
||||
[self.iou_token.weight, self.mask_tokens.weight], dim=0
|
||||
)
|
||||
output_tokens = output_tokens.unsqueeze(0).expand(
|
||||
sparse_prompt_embeddings.size(0), -1, -1
|
||||
)
|
||||
tokens = torch.cat((output_tokens, sparse_prompt_embeddings), dim=1)
|
||||
|
||||
# Expand per-image data in batch direction to be per-mask
|
||||
if repeat_image:
|
||||
src = torch.repeat_interleave(image_embeddings, tokens.shape[0], dim=0)
|
||||
else:
|
||||
assert image_embeddings.shape[0] == tokens.shape[0]
|
||||
src = image_embeddings
|
||||
src = src + dense_prompt_embeddings
|
||||
assert (
|
||||
image_pe.size(0) == 1
|
||||
), "image_pe should have size 1 in batch dim (from `get_dense_pe()`)"
|
||||
pos_src = torch.repeat_interleave(image_pe, tokens.shape[0], dim=0)
|
||||
b, c, h, w = src.shape
|
||||
|
||||
# Run the transformer
|
||||
hs, src = self.transformer(src, pos_src, tokens)
|
||||
iou_token_out = hs[:, s, :]
|
||||
mask_tokens_out = hs[:, s + 1 : (s + 1 + self.num_mask_tokens), :]
|
||||
|
||||
# Upscale mask embeddings and predict masks using the mask tokens
|
||||
src = src.transpose(1, 2).view(b, c, h, w)
|
||||
if not self.use_high_res_features:
|
||||
upscaled_embedding = self.output_upscaling(src)
|
||||
else:
|
||||
dc1, ln1, act1, dc2, act2 = self.output_upscaling
|
||||
feat_s0, feat_s1 = high_res_features
|
||||
upscaled_embedding = act1(ln1(dc1(src) + feat_s1))
|
||||
upscaled_embedding = act2(dc2(upscaled_embedding) + feat_s0)
|
||||
|
||||
hyper_in_list: List[torch.Tensor] = []
|
||||
for i in range(self.num_mask_tokens):
|
||||
hyper_in_list.append(
|
||||
self.output_hypernetworks_mlps[i](mask_tokens_out[:, i, :])
|
||||
)
|
||||
hyper_in = torch.stack(hyper_in_list, dim=1)
|
||||
b, c, h, w = upscaled_embedding.shape
|
||||
masks = (hyper_in @ upscaled_embedding.view(b, c, h * w)).view(b, -1, h, w)
|
||||
|
||||
# Generate mask quality predictions
|
||||
iou_pred = self.iou_prediction_head(iou_token_out)
|
||||
if self.pred_obj_scores:
|
||||
assert s == 1
|
||||
object_score_logits = self.pred_obj_score_head(hs[:, 0, :])
|
||||
else:
|
||||
# Obj scores logits - default to 10.0, i.e. assuming the object is present, sigmoid(10)=1
|
||||
object_score_logits = 10.0 * iou_pred.new_ones(iou_pred.shape[0], 1)
|
||||
|
||||
return masks, iou_pred, mask_tokens_out, object_score_logits
|
||||
|
||||
def _get_stability_scores(self, mask_logits):
|
||||
"""
|
||||
Compute stability scores of the mask logits based on the IoU between upper and
|
||||
lower thresholds, similar to https://github.com/fairinternal/onevision/pull/568.
|
||||
"""
|
||||
mask_logits = mask_logits.flatten(-2)
|
||||
stability_delta = self.dynamic_multimask_stability_delta
|
||||
area_i = torch.sum(mask_logits > stability_delta, dim=-1).float()
|
||||
area_u = torch.sum(mask_logits > -stability_delta, dim=-1).float()
|
||||
stability_scores = torch.where(area_u > 0, area_i / area_u, 1.0)
|
||||
return stability_scores
|
||||
|
||||
def _dynamic_multimask_via_stability(self, all_mask_logits, all_iou_scores):
|
||||
"""
|
||||
When outputting a single mask, if the stability score from the current single-mask
|
||||
output (based on output token 0) falls below a threshold, we instead select from
|
||||
multi-mask outputs (based on output token 1~3) the mask with the highest predicted
|
||||
IoU score. This is intended to ensure a valid mask for both clicking and tracking.
|
||||
"""
|
||||
# The best mask from multimask output tokens (1~3)
|
||||
multimask_logits = all_mask_logits[:, 1:, :, :]
|
||||
multimask_iou_scores = all_iou_scores[:, 1:]
|
||||
best_scores_inds = torch.argmax(multimask_iou_scores, dim=-1)
|
||||
batch_inds = torch.arange(
|
||||
multimask_iou_scores.size(0), device=all_iou_scores.device
|
||||
)
|
||||
best_multimask_logits = multimask_logits[batch_inds, best_scores_inds]
|
||||
best_multimask_logits = best_multimask_logits.unsqueeze(1)
|
||||
best_multimask_iou_scores = multimask_iou_scores[batch_inds, best_scores_inds]
|
||||
best_multimask_iou_scores = best_multimask_iou_scores.unsqueeze(1)
|
||||
|
||||
# The mask from singlemask output token 0 and its stability score
|
||||
singlemask_logits = all_mask_logits[:, 0:1, :, :]
|
||||
singlemask_iou_scores = all_iou_scores[:, 0:1]
|
||||
stability_scores = self._get_stability_scores(singlemask_logits)
|
||||
is_stable = stability_scores >= self.dynamic_multimask_stability_thresh
|
||||
|
||||
# Dynamically fall back to best multimask output upon low stability scores.
|
||||
mask_logits_out = torch.where(
|
||||
is_stable[..., None, None].expand_as(singlemask_logits),
|
||||
singlemask_logits,
|
||||
best_multimask_logits,
|
||||
)
|
||||
iou_scores_out = torch.where(
|
||||
is_stable.expand_as(singlemask_iou_scores),
|
||||
singlemask_iou_scores,
|
||||
best_multimask_iou_scores,
|
||||
)
|
||||
return mask_logits_out, iou_scores_out
|
||||
@@ -0,0 +1,241 @@
|
||||
# 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.
|
||||
|
||||
from typing import Any, Optional, Tuple, Type
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch import nn
|
||||
# from model.segment_anything_2.sam2.modeling.position_encoding import PositionEmbeddingRandom
|
||||
|
||||
from model.segment_anything_2.sam2.modeling.sam2_utils import LayerNorm2d
|
||||
|
||||
|
||||
class PromptEncoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
embed_dim: int,
|
||||
image_embedding_size: Tuple[int, int],
|
||||
input_image_size: Tuple[int, int],
|
||||
mask_in_chans: int,
|
||||
activation: Type[nn.Module] = nn.GELU,
|
||||
) -> None:
|
||||
"""
|
||||
Encodes prompts for input to SAM's mask decoder.
|
||||
|
||||
Arguments:
|
||||
embed_dim (int): The prompts' embedding dimension
|
||||
image_embedding_size (tuple(int, int)): The spatial size of the
|
||||
image embedding, as (H, W).
|
||||
input_image_size (int): The padded size of the image as input
|
||||
to the image encoder, as (H, W).
|
||||
mask_in_chans (int): The number of hidden channels used for
|
||||
encoding input masks.
|
||||
activation (nn.Module): The activation to use when encoding
|
||||
input masks.
|
||||
"""
|
||||
super().__init__()
|
||||
self.embed_dim = embed_dim
|
||||
self.input_image_size = input_image_size
|
||||
self.image_embedding_size = image_embedding_size
|
||||
self.pe_layer = PositionEmbeddingRandom(embed_dim // 2)
|
||||
|
||||
self.num_point_embeddings: int = 4 # pos/neg point + 2 box corners
|
||||
point_embeddings = [
|
||||
nn.Embedding(1, embed_dim) for i in range(self.num_point_embeddings)
|
||||
]
|
||||
self.point_embeddings = nn.ModuleList(point_embeddings)
|
||||
self.not_a_point_embed = nn.Embedding(1, embed_dim)
|
||||
|
||||
self.mask_input_size = (
|
||||
4 * image_embedding_size[0],
|
||||
4 * image_embedding_size[1],
|
||||
)
|
||||
self.mask_downscaling = nn.Sequential(
|
||||
nn.Conv2d(1, mask_in_chans // 4, kernel_size=2, stride=2),
|
||||
LayerNorm2d(mask_in_chans // 4),
|
||||
activation(),
|
||||
nn.Conv2d(mask_in_chans // 4, mask_in_chans, kernel_size=2, stride=2),
|
||||
LayerNorm2d(mask_in_chans),
|
||||
activation(),
|
||||
nn.Conv2d(mask_in_chans, embed_dim, kernel_size=1),
|
||||
)
|
||||
self.no_mask_embed = nn.Embedding(1, embed_dim)
|
||||
|
||||
def get_dense_pe(self) -> torch.Tensor:
|
||||
"""
|
||||
Returns the positional encoding used to encode point prompts,
|
||||
applied to a dense set of points the shape of the image encoding.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Positional encoding with shape
|
||||
1x(embed_dim)x(embedding_h)x(embedding_w)
|
||||
"""
|
||||
return self.pe_layer(self.image_embedding_size).unsqueeze(0)
|
||||
|
||||
def _embed_points(
|
||||
self,
|
||||
points: torch.Tensor,
|
||||
labels: torch.Tensor,
|
||||
pad: bool,
|
||||
) -> torch.Tensor:
|
||||
"""Embeds point prompts."""
|
||||
points = points + 0.5 # Shift to center of pixel
|
||||
if pad:
|
||||
padding_point = torch.zeros((points.shape[0], 1, 2), device=points.device)
|
||||
padding_label = -torch.ones((labels.shape[0], 1), device=labels.device)
|
||||
points = torch.cat([points, padding_point], dim=1)
|
||||
labels = torch.cat([labels, padding_label], dim=1)
|
||||
point_embedding = self.pe_layer.forward_with_coords(
|
||||
points, self.input_image_size
|
||||
)
|
||||
point_embedding[labels == -1] = 0.0
|
||||
point_embedding[labels == -1] += self.not_a_point_embed.weight
|
||||
point_embedding[labels == 0] += self.point_embeddings[0].weight
|
||||
point_embedding[labels == 1] += self.point_embeddings[1].weight
|
||||
point_embedding[labels == 2] += self.point_embeddings[2].weight
|
||||
point_embedding[labels == 3] += self.point_embeddings[3].weight
|
||||
return point_embedding
|
||||
|
||||
def _embed_boxes(self, boxes: torch.Tensor) -> torch.Tensor:
|
||||
"""Embeds box prompts."""
|
||||
boxes = boxes + 0.5 # Shift to center of pixel
|
||||
coords = boxes.reshape(-1, 2, 2)
|
||||
corner_embedding = self.pe_layer.forward_with_coords(
|
||||
coords, self.input_image_size
|
||||
)
|
||||
corner_embedding[:, 0, :] += self.point_embeddings[2].weight
|
||||
corner_embedding[:, 1, :] += self.point_embeddings[3].weight
|
||||
return corner_embedding
|
||||
|
||||
def _embed_masks(self, masks: torch.Tensor) -> torch.Tensor:
|
||||
"""Embeds mask inputs."""
|
||||
mask_embedding = self.mask_downscaling(masks)
|
||||
return mask_embedding
|
||||
|
||||
def _get_batch_size(
|
||||
self,
|
||||
points: Optional[Tuple[torch.Tensor, torch.Tensor]],
|
||||
boxes: Optional[torch.Tensor],
|
||||
masks: Optional[torch.Tensor],
|
||||
text_embeds: Optional[torch.Tensor],
|
||||
) -> int:
|
||||
"""
|
||||
Gets the batch size of the output given the batch size of the input prompts.
|
||||
"""
|
||||
if points is not None:
|
||||
return points[0].shape[0]
|
||||
elif boxes is not None:
|
||||
return boxes.shape[0]
|
||||
elif masks is not None:
|
||||
return masks.shape[0]
|
||||
elif text_embeds is not None:
|
||||
return text_embeds.shape[0]
|
||||
else:
|
||||
return 1
|
||||
|
||||
def _get_device(self) -> torch.device:
|
||||
return self.point_embeddings[0].weight.device
|
||||
|
||||
def forward(
|
||||
self,
|
||||
points: Optional[Tuple[torch.Tensor, torch.Tensor]],
|
||||
boxes: Optional[torch.Tensor],
|
||||
masks: Optional[torch.Tensor],
|
||||
text_embeds: Optional[torch.Tensor],
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Embeds different types of prompts, returning both sparse and dense
|
||||
embeddings.
|
||||
|
||||
Arguments:
|
||||
points (tuple(torch.Tensor, torch.Tensor) or none): point coordinates
|
||||
and labels to embed.
|
||||
boxes (torch.Tensor or none): boxes to embed
|
||||
masks (torch.Tensor or none): masks to embed
|
||||
|
||||
Returns:
|
||||
torch.Tensor: sparse embeddings for the points and boxes, with shape
|
||||
BxNx(embed_dim), where N is determined by the number of input points
|
||||
and boxes.
|
||||
torch.Tensor: dense embeddings for the masks, in the shape
|
||||
Bx(embed_dim)x(embed_H)x(embed_W)
|
||||
"""
|
||||
bs = self._get_batch_size(points, boxes, masks, text_embeds)
|
||||
sparse_embeddings = torch.empty(
|
||||
(bs, 0, self.embed_dim), device=self._get_device()
|
||||
)
|
||||
if points is not None:
|
||||
coords, labels = points
|
||||
point_embeddings = self._embed_points(coords, labels, pad=(boxes is None))
|
||||
sparse_embeddings = torch.cat([sparse_embeddings, point_embeddings], dim=1)
|
||||
if boxes is not None:
|
||||
box_embeddings = self._embed_boxes(boxes)
|
||||
sparse_embeddings = torch.cat([sparse_embeddings, box_embeddings], dim=1)
|
||||
|
||||
if text_embeds is not None:
|
||||
sparse_embeddings = torch.cat([sparse_embeddings, text_embeds], dim=1)
|
||||
|
||||
if masks is not None:
|
||||
dense_embeddings = self._embed_masks(masks)
|
||||
else:
|
||||
dense_embeddings = self.no_mask_embed.weight.reshape(1, -1, 1, 1).expand(
|
||||
bs, -1, self.image_embedding_size[0], self.image_embedding_size[1]
|
||||
)
|
||||
|
||||
return sparse_embeddings, dense_embeddings
|
||||
|
||||
|
||||
class PositionEmbeddingRandom(nn.Module):
|
||||
"""
|
||||
Positional encoding using random spatial frequencies.
|
||||
"""
|
||||
|
||||
def __init__(self, num_pos_feats: int = 64, scale: Optional[float] = None) -> None:
|
||||
super().__init__()
|
||||
if scale is None or scale <= 0.0:
|
||||
scale = 1.0
|
||||
self.register_buffer(
|
||||
"positional_encoding_gaussian_matrix",
|
||||
scale * torch.randn((2, num_pos_feats)),
|
||||
)
|
||||
|
||||
def _pe_encoding(self, coords: torch.Tensor) -> torch.Tensor:
|
||||
"""Positionally encode points that are normalized to [0,1]."""
|
||||
# assuming coords are in [0, 1]^2 square and have d_1 x ... x d_n x 2 shape
|
||||
coords = 2 * coords - 1
|
||||
|
||||
if coords.dtype != self.positional_encoding_gaussian_matrix.dtype:
|
||||
coords = coords.to(self.positional_encoding_gaussian_matrix.dtype)
|
||||
|
||||
coords = coords @ self.positional_encoding_gaussian_matrix
|
||||
coords = 2 * np.pi * coords
|
||||
# outputs d_1 x ... x d_n x C shape
|
||||
return torch.cat([torch.sin(coords), torch.cos(coords)], dim=-1)
|
||||
|
||||
def forward(self, size: Tuple[int, int]) -> torch.Tensor:
|
||||
"""Generate positional encoding for a grid of the specified size."""
|
||||
h, w = size
|
||||
device: Any = self.positional_encoding_gaussian_matrix.device
|
||||
grid = torch.ones(
|
||||
(h, w), device=device, dtype=self.positional_encoding_gaussian_matrix.dtype
|
||||
)
|
||||
y_embed = grid.cumsum(dim=0) - 0.5
|
||||
x_embed = grid.cumsum(dim=1) - 0.5
|
||||
y_embed = y_embed / h
|
||||
x_embed = x_embed / w
|
||||
|
||||
pe = self._pe_encoding(torch.stack([x_embed, y_embed], dim=-1))
|
||||
return pe.permute(2, 0, 1) # C x H x W
|
||||
|
||||
def forward_with_coords(
|
||||
self, coords_input: torch.Tensor, image_size: Tuple[int, int]
|
||||
) -> torch.Tensor:
|
||||
"""Positionally encode points that are not normalized to [0,1]."""
|
||||
coords = coords_input.clone()
|
||||
coords[:, :, 0] = coords[:, :, 0] / image_size[1]
|
||||
coords[:, :, 1] = coords[:, :, 1] / image_size[0]
|
||||
return self._pe_encoding(coords.to(torch.float)) # B x N x C
|
||||
@@ -0,0 +1,330 @@
|
||||
# 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 math
|
||||
import warnings
|
||||
from functools import partial
|
||||
from typing import Tuple, Type
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn, Tensor
|
||||
|
||||
from model.segment_anything_2.sam2.modeling.position_encoding import apply_rotary_enc, compute_axial_cis
|
||||
|
||||
from model.segment_anything_2.sam2.modeling.sam2_utils import MLP
|
||||
from model.segment_anything_2.sam2.utils.misc import get_sdpa_settings
|
||||
|
||||
warnings.simplefilter(action="ignore", category=FutureWarning)
|
||||
# OLD_GPU, USE_FLASH_ATTN, MATH_KERNEL_ON = get_sdpa_settings()
|
||||
USE_FLASH_ATTN = False
|
||||
MATH_KERNEL_ON = True
|
||||
OLD_GPU = True
|
||||
|
||||
|
||||
class TwoWayTransformer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
depth: int,
|
||||
embedding_dim: int,
|
||||
num_heads: int,
|
||||
mlp_dim: int,
|
||||
activation: Type[nn.Module] = nn.ReLU,
|
||||
attention_downsample_rate: int = 2,
|
||||
) -> None:
|
||||
"""
|
||||
A transformer decoder that attends to an input image using
|
||||
queries whose positional embedding is supplied.
|
||||
|
||||
Args:
|
||||
depth (int): number of layers in the transformer
|
||||
embedding_dim (int): the channel dimension for the input embeddings
|
||||
num_heads (int): the number of heads for multihead attention. Must
|
||||
divide embedding_dim
|
||||
mlp_dim (int): the channel dimension internal to the MLP block
|
||||
activation (nn.Module): the activation to use in the MLP block
|
||||
"""
|
||||
super().__init__()
|
||||
self.depth = depth
|
||||
self.embedding_dim = embedding_dim
|
||||
self.num_heads = num_heads
|
||||
self.mlp_dim = mlp_dim
|
||||
self.layers = nn.ModuleList()
|
||||
|
||||
for i in range(depth):
|
||||
self.layers.append(
|
||||
TwoWayAttentionBlock(
|
||||
embedding_dim=embedding_dim,
|
||||
num_heads=num_heads,
|
||||
mlp_dim=mlp_dim,
|
||||
activation=activation,
|
||||
attention_downsample_rate=attention_downsample_rate,
|
||||
skip_first_layer_pe=(i == 0),
|
||||
)
|
||||
)
|
||||
|
||||
self.final_attn_token_to_image = Attention(
|
||||
embedding_dim, num_heads, downsample_rate=attention_downsample_rate
|
||||
)
|
||||
self.norm_final_attn = nn.LayerNorm(embedding_dim)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
image_embedding: Tensor,
|
||||
image_pe: Tensor,
|
||||
point_embedding: Tensor,
|
||||
) -> Tuple[Tensor, Tensor]:
|
||||
"""
|
||||
Args:
|
||||
image_embedding (torch.Tensor): image to attend to. Should be shape
|
||||
B x embedding_dim x h x w for any h and w.
|
||||
image_pe (torch.Tensor): the positional encoding to add to the image. Must
|
||||
have the same shape as image_embedding.
|
||||
point_embedding (torch.Tensor): the embedding to add to the query points.
|
||||
Must have shape B x N_points x embedding_dim for any N_points.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: the processed point_embedding
|
||||
torch.Tensor: the processed image_embedding
|
||||
"""
|
||||
# BxCxHxW -> BxHWxC == B x N_image_tokens x C
|
||||
bs, c, h, w = image_embedding.shape
|
||||
image_embedding = image_embedding.flatten(2).permute(0, 2, 1)
|
||||
image_pe = image_pe.flatten(2).permute(0, 2, 1)
|
||||
|
||||
# Prepare queries
|
||||
queries = point_embedding
|
||||
keys = image_embedding
|
||||
|
||||
# Apply transformer blocks and final layernorm
|
||||
for layer in self.layers:
|
||||
queries, keys = layer(
|
||||
queries=queries,
|
||||
keys=keys,
|
||||
query_pe=point_embedding,
|
||||
key_pe=image_pe,
|
||||
)
|
||||
|
||||
# Apply the final attention layer from the points to the image
|
||||
q = queries + point_embedding
|
||||
k = keys + image_pe
|
||||
attn_out = self.final_attn_token_to_image(q=q, k=k, v=keys)
|
||||
queries = queries + attn_out
|
||||
queries = self.norm_final_attn(queries)
|
||||
|
||||
return queries, keys
|
||||
|
||||
|
||||
class TwoWayAttentionBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
embedding_dim: int,
|
||||
num_heads: int,
|
||||
mlp_dim: int = 2048,
|
||||
activation: Type[nn.Module] = nn.ReLU,
|
||||
attention_downsample_rate: int = 2,
|
||||
skip_first_layer_pe: bool = False,
|
||||
) -> None:
|
||||
"""
|
||||
A transformer block with four layers: (1) self-attention of sparse
|
||||
inputs, (2) cross attention of sparse inputs to dense inputs, (3) mlp
|
||||
block on sparse inputs, and (4) cross attention of dense inputs to sparse
|
||||
inputs.
|
||||
|
||||
Arguments:
|
||||
embedding_dim (int): the channel dimension of the embeddings
|
||||
num_heads (int): the number of heads in the attention layers
|
||||
mlp_dim (int): the hidden dimension of the mlp block
|
||||
activation (nn.Module): the activation of the mlp block
|
||||
skip_first_layer_pe (bool): skip the PE on the first layer
|
||||
"""
|
||||
super().__init__()
|
||||
self.self_attn = Attention(embedding_dim, num_heads)
|
||||
self.norm1 = nn.LayerNorm(embedding_dim)
|
||||
|
||||
self.cross_attn_token_to_image = Attention(
|
||||
embedding_dim, num_heads, downsample_rate=attention_downsample_rate
|
||||
)
|
||||
self.norm2 = nn.LayerNorm(embedding_dim)
|
||||
|
||||
self.mlp = MLP(
|
||||
embedding_dim, mlp_dim, embedding_dim, num_layers=2, activation=activation
|
||||
)
|
||||
self.norm3 = nn.LayerNorm(embedding_dim)
|
||||
|
||||
self.norm4 = nn.LayerNorm(embedding_dim)
|
||||
self.cross_attn_image_to_token = Attention(
|
||||
embedding_dim, num_heads, downsample_rate=attention_downsample_rate
|
||||
)
|
||||
|
||||
self.skip_first_layer_pe = skip_first_layer_pe
|
||||
|
||||
def forward(
|
||||
self, queries: Tensor, keys: Tensor, query_pe: Tensor, key_pe: Tensor
|
||||
) -> Tuple[Tensor, Tensor]:
|
||||
# Self attention block
|
||||
if self.skip_first_layer_pe:
|
||||
queries = self.self_attn(q=queries, k=queries, v=queries)
|
||||
else:
|
||||
q = queries + query_pe
|
||||
attn_out = self.self_attn(q=q, k=q, v=queries)
|
||||
queries = queries + attn_out
|
||||
queries = self.norm1(queries)
|
||||
|
||||
# Cross attention block, tokens attending to image embedding
|
||||
q = queries + query_pe
|
||||
k = keys + key_pe
|
||||
attn_out = self.cross_attn_token_to_image(q=q, k=k, v=keys)
|
||||
queries = queries + attn_out
|
||||
queries = self.norm2(queries)
|
||||
|
||||
# MLP block
|
||||
mlp_out = self.mlp(queries)
|
||||
queries = queries + mlp_out
|
||||
queries = self.norm3(queries)
|
||||
|
||||
# Cross attention block, image embedding attending to tokens
|
||||
q = queries + query_pe
|
||||
k = keys + key_pe
|
||||
attn_out = self.cross_attn_image_to_token(q=k, k=q, v=queries)
|
||||
keys = keys + attn_out
|
||||
keys = self.norm4(keys)
|
||||
|
||||
return queries, keys
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
"""
|
||||
An attention layer that allows for downscaling the size of the embedding
|
||||
after projection to queries, keys, and values.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
embedding_dim: int,
|
||||
num_heads: int,
|
||||
downsample_rate: int = 1,
|
||||
dropout: float = 0.0,
|
||||
kv_in_dim: int = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.embedding_dim = embedding_dim
|
||||
self.kv_in_dim = kv_in_dim if kv_in_dim is not None else embedding_dim
|
||||
self.internal_dim = embedding_dim // downsample_rate
|
||||
self.num_heads = num_heads
|
||||
assert (
|
||||
self.internal_dim % num_heads == 0
|
||||
), "num_heads must divide embedding_dim."
|
||||
|
||||
self.q_proj = nn.Linear(embedding_dim, self.internal_dim)
|
||||
self.k_proj = nn.Linear(self.kv_in_dim, self.internal_dim)
|
||||
self.v_proj = nn.Linear(self.kv_in_dim, self.internal_dim)
|
||||
self.out_proj = nn.Linear(self.internal_dim, embedding_dim)
|
||||
|
||||
self.dropout_p = dropout
|
||||
|
||||
def _separate_heads(self, x: Tensor, num_heads: int) -> Tensor:
|
||||
b, n, c = x.shape
|
||||
x = x.reshape(b, n, num_heads, c // num_heads)
|
||||
return x.transpose(1, 2) # B x N_heads x N_tokens x C_per_head
|
||||
|
||||
def _recombine_heads(self, x: Tensor) -> Tensor:
|
||||
b, n_heads, n_tokens, c_per_head = x.shape
|
||||
x = x.transpose(1, 2)
|
||||
return x.reshape(b, n_tokens, n_heads * c_per_head) # B x N_tokens x C
|
||||
|
||||
def forward(self, q: Tensor, k: Tensor, v: Tensor) -> Tensor:
|
||||
# Input projections
|
||||
q = self.q_proj(q)
|
||||
k = self.k_proj(k)
|
||||
v = self.v_proj(v)
|
||||
|
||||
# Separate into heads
|
||||
q = self._separate_heads(q, self.num_heads)
|
||||
k = self._separate_heads(k, self.num_heads)
|
||||
v = self._separate_heads(v, self.num_heads)
|
||||
|
||||
dropout_p = self.dropout_p if self.training else 0.0
|
||||
# Attention
|
||||
with torch.backends.cuda.sdp_kernel(
|
||||
enable_flash=USE_FLASH_ATTN,
|
||||
# if Flash attention kernel is off, then math kernel needs to be enabled
|
||||
enable_math=(OLD_GPU and dropout_p > 0.0) or MATH_KERNEL_ON,
|
||||
enable_mem_efficient=OLD_GPU,
|
||||
):
|
||||
out = F.scaled_dot_product_attention(q, k, v, dropout_p=dropout_p)
|
||||
|
||||
out = self._recombine_heads(out)
|
||||
out = self.out_proj(out)
|
||||
|
||||
return out
|
||||
|
||||
|
||||
class RoPEAttention(Attention):
|
||||
"""Attention with rotary position encoding."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*args,
|
||||
rope_theta=10000.0,
|
||||
# whether to repeat q rope to match k length
|
||||
# this is needed for cross-attention to memories
|
||||
rope_k_repeat=False,
|
||||
feat_sizes=(32, 32), # [w, h] for stride 16 feats at 512 resolution
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
self.compute_cis = partial(
|
||||
compute_axial_cis, dim=self.internal_dim // self.num_heads, theta=rope_theta
|
||||
)
|
||||
freqs_cis = self.compute_cis(end_x=feat_sizes[0], end_y=feat_sizes[1])
|
||||
self.freqs_cis = freqs_cis
|
||||
self.rope_k_repeat = rope_k_repeat
|
||||
|
||||
def forward(
|
||||
self, q: Tensor, k: Tensor, v: Tensor, num_k_exclude_rope: int = 0
|
||||
) -> Tensor:
|
||||
# Input projections
|
||||
q = self.q_proj(q)
|
||||
k = self.k_proj(k)
|
||||
v = self.v_proj(v)
|
||||
|
||||
# Separate into heads
|
||||
q = self._separate_heads(q, self.num_heads)
|
||||
k = self._separate_heads(k, self.num_heads)
|
||||
v = self._separate_heads(v, self.num_heads)
|
||||
|
||||
# Apply rotary position encoding
|
||||
w = h = math.sqrt(q.shape[-2])
|
||||
self.freqs_cis = self.freqs_cis.to(q.device)
|
||||
if self.freqs_cis.shape[0] != q.shape[-2]:
|
||||
self.freqs_cis = self.compute_cis(end_x=w, end_y=h).to(q.device)
|
||||
if q.shape[-2] != k.shape[-2]:
|
||||
assert self.rope_k_repeat
|
||||
|
||||
num_k_rope = k.size(-2) - num_k_exclude_rope
|
||||
q, k[:, :, :num_k_rope] = apply_rotary_enc(
|
||||
q,
|
||||
k[:, :, :num_k_rope],
|
||||
freqs_cis=self.freqs_cis,
|
||||
repeat_freqs_k=self.rope_k_repeat,
|
||||
)
|
||||
|
||||
dropout_p = self.dropout_p if self.training else 0.0
|
||||
# Attention
|
||||
with torch.backends.cuda.sdp_kernel(
|
||||
enable_flash=USE_FLASH_ATTN,
|
||||
# if Flash attention kernel is off, then math kernel needs to be enabled
|
||||
enable_math=(OLD_GPU and dropout_p > 0.0) or MATH_KERNEL_ON,
|
||||
enable_mem_efficient=OLD_GPU,
|
||||
):
|
||||
out = F.scaled_dot_product_attention(q, k, v, dropout_p=dropout_p)
|
||||
|
||||
out = self._recombine_heads(out)
|
||||
out = self.out_proj(out)
|
||||
|
||||
return out
|
||||
@@ -0,0 +1,833 @@
|
||||
# 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 torch
|
||||
import torch.distributed
|
||||
import torch.nn.functional as F
|
||||
|
||||
from torch.nn.init import trunc_normal_
|
||||
|
||||
from model.segment_anything_2.sam2.modeling.sam.mask_decoder import MaskDecoder
|
||||
from model.segment_anything_2.sam2.modeling.sam.prompt_encoder import PromptEncoder
|
||||
from model.segment_anything_2.sam2.modeling.sam.transformer import TwoWayTransformer
|
||||
from model.segment_anything_2.sam2.modeling.sam2_utils import get_1d_sine_pe, MLP, select_closest_cond_frames
|
||||
|
||||
# a large negative value as a placeholder score for missing objects
|
||||
NO_OBJ_SCORE = -1024.0
|
||||
|
||||
|
||||
class SAM2Base(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
image_encoder,
|
||||
memory_attention,
|
||||
memory_encoder,
|
||||
num_maskmem=7, # default 1 input frame + 6 previous frames
|
||||
image_size=512,
|
||||
backbone_stride=16, # stride of the image backbone output
|
||||
sigmoid_scale_for_mem_enc=1.0, # scale factor for mask sigmoid prob
|
||||
sigmoid_bias_for_mem_enc=0.0, # bias factor for mask sigmoid prob
|
||||
# During evaluation, whether to binarize the sigmoid mask logits on interacted frames with clicks
|
||||
binarize_mask_from_pts_for_mem_enc=False,
|
||||
use_mask_input_as_output_without_sam=False, # on frames with mask input, whether to directly output the input mask without using a SAM prompt encoder + mask decoder
|
||||
# The maximum number of conditioning frames to participate in the memory attention (-1 means no limit; if there are more conditioning frames than this limit,
|
||||
# we only cross-attend to the temporally closest `max_cond_frames_in_attn` conditioning frames in the encoder when tracking each frame). This gives the model
|
||||
# a temporal locality when handling a large number of annotated frames (since closer frames should be more important) and also avoids GPU OOM.
|
||||
max_cond_frames_in_attn=-1,
|
||||
# on the first frame, whether to directly add the no-memory embedding to the image feature
|
||||
# (instead of using the transformer encoder)
|
||||
directly_add_no_mem_embed=False,
|
||||
# whether to use high-resolution feature maps in the SAM mask decoder
|
||||
use_high_res_features_in_sam=False,
|
||||
# whether to output multiple (3) masks for the first click on initial conditioning frames
|
||||
multimask_output_in_sam=False,
|
||||
# the minimum and maximum number of clicks to use multimask_output_in_sam (only relevant when `multimask_output_in_sam=True`;
|
||||
# default is 1 for both, meaning that only the first click gives multimask output; also note that a box counts as two points)
|
||||
multimask_min_pt_num=1,
|
||||
multimask_max_pt_num=1,
|
||||
# whether to also use multimask output for tracking (not just for the first click on initial conditioning frames; only relevant when `multimask_output_in_sam=True`)
|
||||
multimask_output_for_tracking=False,
|
||||
# Whether to use multimask tokens for obj ptr; Only relevant when both
|
||||
# use_obj_ptrs_in_encoder=True and multimask_output_for_tracking=True
|
||||
use_multimask_token_for_obj_ptr: bool = False,
|
||||
# whether to use sigmoid to restrict ious prediction to [0-1]
|
||||
iou_prediction_use_sigmoid=False,
|
||||
# The memory bank's temporal stride during evaluation (i.e. the `r` parameter in XMem and Cutie; XMem and Cutie use r=5).
|
||||
# For r>1, the (self.num_maskmem - 1) non-conditioning memory frames consist of
|
||||
# (self.num_maskmem - 2) nearest frames from every r-th frames, plus the last frame.
|
||||
memory_temporal_stride_for_eval=1,
|
||||
# if `add_all_frames_to_correct_as_cond` is True, we also append to the conditioning frame list any frame that receives a later correction click
|
||||
# if `add_all_frames_to_correct_as_cond` is False, we conditioning frame list to only use those initial conditioning frames
|
||||
add_all_frames_to_correct_as_cond=False,
|
||||
# whether to apply non-overlapping constraints on the object masks in the memory encoder during evaluation (to avoid/alleviate superposing masks)
|
||||
non_overlap_masks_for_mem_enc=False,
|
||||
# whether to cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
|
||||
use_obj_ptrs_in_encoder=False,
|
||||
# the maximum number of object pointers from other frames in encoder cross attention (only relevant when `use_obj_ptrs_in_encoder=True`)
|
||||
max_obj_ptrs_in_encoder=16,
|
||||
# whether to add temporal positional encoding to the object pointers in the encoder (only relevant when `use_obj_ptrs_in_encoder=True`)
|
||||
add_tpos_enc_to_obj_ptrs=True,
|
||||
# whether to add an extra linear projection layer for the temporal positional encoding in the object pointers to avoid potential interference
|
||||
# with spatial positional encoding (only relevant when both `use_obj_ptrs_in_encoder=True` and `add_tpos_enc_to_obj_ptrs=True`)
|
||||
proj_tpos_enc_in_obj_ptrs=False,
|
||||
# whether to only attend to object pointers in the past (before the current frame) in the encoder during evaluation
|
||||
# (only relevant when `use_obj_ptrs_in_encoder=True`; this might avoid pointer information too far in the future to distract the initial tracking)
|
||||
only_obj_ptrs_in_the_past_for_eval=False,
|
||||
# Whether to predict if there is an object in the frame
|
||||
pred_obj_scores: bool = False,
|
||||
# Whether to use an MLP to predict object scores
|
||||
pred_obj_scores_mlp: bool = False,
|
||||
# Only relevant if pred_obj_scores=True and use_obj_ptrs_in_encoder=True;
|
||||
# Whether to have a fixed no obj pointer when there is no object present
|
||||
# or to use it as an additive embedding with obj_ptr produced by decoder
|
||||
fixed_no_obj_ptr: bool = False,
|
||||
# Soft no object, i.e. mix in no_obj_ptr softly,
|
||||
# hope to make recovery easier if there is a mistake and mitigate accumulation of errors
|
||||
soft_no_obj_ptr: bool = False,
|
||||
use_mlp_for_obj_ptr_proj: bool = False,
|
||||
# extra arguments used to construct the SAM mask decoder; if not None, it should be a dict of kwargs to be passed into `MaskDecoder` class.
|
||||
sam_mask_decoder_extra_args=None,
|
||||
compile_image_encoder: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
# Part 1: the image backbone
|
||||
self.image_encoder = image_encoder
|
||||
# Use level 0, 1, 2 for high-res setting, or just level 2 for the default setting
|
||||
self.use_high_res_features_in_sam = use_high_res_features_in_sam
|
||||
self.num_feature_levels = 3 if use_high_res_features_in_sam else 1
|
||||
self.use_obj_ptrs_in_encoder = use_obj_ptrs_in_encoder
|
||||
self.max_obj_ptrs_in_encoder = max_obj_ptrs_in_encoder
|
||||
if use_obj_ptrs_in_encoder:
|
||||
# A conv layer to downsample the mask prompt to stride 4 (the same stride as
|
||||
# low-res SAM mask logits) and to change its scales from 0~1 to SAM logit scale,
|
||||
# so that it can be fed into the SAM mask decoder to generate a pointer.
|
||||
self.mask_downsample = torch.nn.Conv2d(1, 1, kernel_size=4, stride=4)
|
||||
self.add_tpos_enc_to_obj_ptrs = add_tpos_enc_to_obj_ptrs
|
||||
if proj_tpos_enc_in_obj_ptrs:
|
||||
assert add_tpos_enc_to_obj_ptrs # these options need to be used together
|
||||
self.proj_tpos_enc_in_obj_ptrs = proj_tpos_enc_in_obj_ptrs
|
||||
self.only_obj_ptrs_in_the_past_for_eval = only_obj_ptrs_in_the_past_for_eval
|
||||
|
||||
# Part 2: memory attention to condition current frame's visual features
|
||||
# with memories (and obj ptrs) from past frames
|
||||
self.memory_attention = memory_attention
|
||||
self.hidden_dim = memory_attention.d_model
|
||||
|
||||
# Part 3: memory encoder for the previous frame's outputs
|
||||
self.memory_encoder = memory_encoder
|
||||
self.mem_dim = self.hidden_dim
|
||||
if hasattr(self.memory_encoder, "out_proj") and hasattr(
|
||||
self.memory_encoder.out_proj, "weight"
|
||||
):
|
||||
# if there is compression of memories along channel dim
|
||||
self.mem_dim = self.memory_encoder.out_proj.weight.shape[0]
|
||||
self.num_maskmem = num_maskmem # Number of memories accessible
|
||||
# Temporal encoding of the memories
|
||||
self.maskmem_tpos_enc = torch.nn.Parameter(
|
||||
torch.zeros(num_maskmem, 1, 1, self.mem_dim)
|
||||
)
|
||||
trunc_normal_(self.maskmem_tpos_enc, std=0.02)
|
||||
# a single token to indicate no memory embedding from previous frames
|
||||
self.no_mem_embed = torch.nn.Parameter(torch.zeros(1, 1, self.hidden_dim))
|
||||
self.no_mem_pos_enc = torch.nn.Parameter(torch.zeros(1, 1, self.hidden_dim))
|
||||
trunc_normal_(self.no_mem_embed, std=0.02)
|
||||
trunc_normal_(self.no_mem_pos_enc, std=0.02)
|
||||
self.directly_add_no_mem_embed = directly_add_no_mem_embed
|
||||
# Apply sigmoid to the output raw mask logits (to turn them from
|
||||
# range (-inf, +inf) to range (0, 1)) before feeding them into the memory encoder
|
||||
self.sigmoid_scale_for_mem_enc = sigmoid_scale_for_mem_enc
|
||||
self.sigmoid_bias_for_mem_enc = sigmoid_bias_for_mem_enc
|
||||
self.binarize_mask_from_pts_for_mem_enc = binarize_mask_from_pts_for_mem_enc
|
||||
self.non_overlap_masks_for_mem_enc = non_overlap_masks_for_mem_enc
|
||||
self.memory_temporal_stride_for_eval = memory_temporal_stride_for_eval
|
||||
# On frames with mask input, whether to directly output the input mask without
|
||||
# using a SAM prompt encoder + mask decoder
|
||||
self.use_mask_input_as_output_without_sam = use_mask_input_as_output_without_sam
|
||||
self.multimask_output_in_sam = multimask_output_in_sam
|
||||
self.multimask_min_pt_num = multimask_min_pt_num
|
||||
self.multimask_max_pt_num = multimask_max_pt_num
|
||||
self.multimask_output_for_tracking = multimask_output_for_tracking
|
||||
self.use_multimask_token_for_obj_ptr = use_multimask_token_for_obj_ptr
|
||||
self.iou_prediction_use_sigmoid = iou_prediction_use_sigmoid
|
||||
|
||||
# Part 4: SAM-style prompt encoder (for both mask and point inputs)
|
||||
# and SAM-style mask decoder for the final mask output
|
||||
self.image_size = image_size
|
||||
self.backbone_stride = backbone_stride
|
||||
self.sam_mask_decoder_extra_args = sam_mask_decoder_extra_args
|
||||
self.pred_obj_scores = pred_obj_scores
|
||||
self.pred_obj_scores_mlp = pred_obj_scores_mlp
|
||||
self.fixed_no_obj_ptr = fixed_no_obj_ptr
|
||||
self.soft_no_obj_ptr = soft_no_obj_ptr
|
||||
if self.fixed_no_obj_ptr:
|
||||
assert self.pred_obj_scores
|
||||
assert self.use_obj_ptrs_in_encoder
|
||||
if self.pred_obj_scores and self.use_obj_ptrs_in_encoder:
|
||||
self.no_obj_ptr = torch.nn.Parameter(torch.zeros(1, self.hidden_dim))
|
||||
trunc_normal_(self.no_obj_ptr, std=0.02)
|
||||
self.use_mlp_for_obj_ptr_proj = use_mlp_for_obj_ptr_proj
|
||||
|
||||
self._build_sam_heads()
|
||||
self.add_all_frames_to_correct_as_cond = add_all_frames_to_correct_as_cond
|
||||
self.max_cond_frames_in_attn = max_cond_frames_in_attn
|
||||
|
||||
# Model compilation
|
||||
if compile_image_encoder:
|
||||
# Compile the forward function (not the full module) to allow loading checkpoints.
|
||||
print(
|
||||
"Image encoder compilation is enabled. First forward pass will be slow."
|
||||
)
|
||||
self.image_encoder.forward = torch.compile(
|
||||
self.image_encoder.forward,
|
||||
mode="max-autotune",
|
||||
fullgraph=True,
|
||||
dynamic=False,
|
||||
)
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
return next(self.parameters()).device
|
||||
|
||||
def forward(self, *args, **kwargs):
|
||||
raise NotImplementedError(
|
||||
"Please use the corresponding methods in SAM2VideoPredictor for inference."
|
||||
"See notebooks/video_predictor_example.ipynb for an example."
|
||||
)
|
||||
|
||||
def _build_sam_heads(self):
|
||||
"""Build SAM-style prompt encoder and mask decoder."""
|
||||
self.sam_prompt_embed_dim = self.hidden_dim
|
||||
self.sam_image_embedding_size = self.image_size // self.backbone_stride
|
||||
|
||||
# build PromptEncoder and MaskDecoder from SAM
|
||||
# (their hyperparameters like `mask_in_chans=16` are from SAM code)
|
||||
self.sam_prompt_encoder = PromptEncoder(
|
||||
embed_dim=self.sam_prompt_embed_dim,
|
||||
image_embedding_size=(
|
||||
self.sam_image_embedding_size,
|
||||
self.sam_image_embedding_size,
|
||||
),
|
||||
input_image_size=(self.image_size, self.image_size),
|
||||
mask_in_chans=16,
|
||||
)
|
||||
self.sam_mask_decoder = MaskDecoder(
|
||||
num_multimask_outputs=3,
|
||||
transformer=TwoWayTransformer(
|
||||
depth=2,
|
||||
embedding_dim=self.sam_prompt_embed_dim,
|
||||
mlp_dim=2048,
|
||||
num_heads=8,
|
||||
),
|
||||
transformer_dim=self.sam_prompt_embed_dim,
|
||||
iou_head_depth=3,
|
||||
iou_head_hidden_dim=256,
|
||||
use_high_res_features=self.use_high_res_features_in_sam,
|
||||
iou_prediction_use_sigmoid=self.iou_prediction_use_sigmoid,
|
||||
pred_obj_scores=self.pred_obj_scores,
|
||||
pred_obj_scores_mlp=self.pred_obj_scores_mlp,
|
||||
use_multimask_token_for_obj_ptr=self.use_multimask_token_for_obj_ptr,
|
||||
**(self.sam_mask_decoder_extra_args or {}),
|
||||
)
|
||||
if self.use_obj_ptrs_in_encoder:
|
||||
# a linear projection on SAM output tokens to turn them into object pointers
|
||||
self.obj_ptr_proj = torch.nn.Linear(self.hidden_dim, self.hidden_dim)
|
||||
if self.use_mlp_for_obj_ptr_proj:
|
||||
self.obj_ptr_proj = MLP(
|
||||
self.hidden_dim, self.hidden_dim, self.hidden_dim, 3
|
||||
)
|
||||
else:
|
||||
self.obj_ptr_proj = torch.nn.Identity()
|
||||
if self.proj_tpos_enc_in_obj_ptrs:
|
||||
# a linear projection on temporal positional encoding in object pointers to
|
||||
# avoid potential interference with spatial positional encoding
|
||||
self.obj_ptr_tpos_proj = torch.nn.Linear(self.hidden_dim, self.mem_dim)
|
||||
else:
|
||||
self.obj_ptr_tpos_proj = torch.nn.Identity()
|
||||
|
||||
def _forward_sam_heads(
|
||||
self,
|
||||
backbone_features,
|
||||
point_inputs=None,
|
||||
mask_inputs=None,
|
||||
text_inputs=None,
|
||||
high_res_features=None,
|
||||
multimask_output=False,
|
||||
):
|
||||
"""
|
||||
Forward SAM prompt encoders and mask heads.
|
||||
|
||||
Inputs:
|
||||
- backbone_features: image features of [B, C, H, W] shape
|
||||
- point_inputs: a dictionary with "point_coords" and "point_labels", where
|
||||
1) "point_coords" has [B, P, 2] shape and float32 dtype and contains the
|
||||
absolute pixel-unit coordinate in (x, y) format of the P input points
|
||||
2) "point_labels" has shape [B, P] and int32 dtype, where 1 means
|
||||
positive clicks, 0 means negative clicks, and -1 means padding
|
||||
- mask_inputs: a mask of [B, 1, H*16, W*16] shape, float or bool, with the
|
||||
same spatial size as the image.
|
||||
- high_res_features: either 1) None or 2) or a list of length 2 containing
|
||||
two feature maps of [B, C, 4*H, 4*W] and [B, C, 2*H, 2*W] shapes respectively,
|
||||
which will be used as high-resolution feature maps for SAM decoder.
|
||||
- multimask_output: if it's True, we output 3 candidate masks and their 3
|
||||
corresponding IoU estimates, and if it's False, we output only 1 mask and
|
||||
its corresponding IoU estimate.
|
||||
|
||||
Outputs:
|
||||
- low_res_multimasks: [B, M, H*4, W*4] shape (where M = 3 if
|
||||
`multimask_output=True` and M = 1 if `multimask_output=False`), the SAM
|
||||
output mask logits (before sigmoid) for the low-resolution masks, with 4x
|
||||
the resolution (1/4 stride) of the input backbone_features.
|
||||
- high_res_multimasks: [B, M, H*16, W*16] shape (where M = 3
|
||||
if `multimask_output=True` and M = 1 if `multimask_output=False`),
|
||||
upsampled from the low-resolution masks, with shape size as the image
|
||||
(stride is 1 pixel).
|
||||
- ious, [B, M] shape, where (where M = 3 if `multimask_output=True` and M = 1
|
||||
if `multimask_output=False`), the estimated IoU of each output mask.
|
||||
- low_res_masks: [B, 1, H*4, W*4] shape, the best mask in `low_res_multimasks`.
|
||||
If `multimask_output=True`, it's the mask with the highest IoU estimate.
|
||||
If `multimask_output=False`, it's the same as `low_res_multimasks`.
|
||||
- high_res_masks: [B, 1, H*16, W*16] shape, the best mask in `high_res_multimasks`.
|
||||
If `multimask_output=True`, it's the mask with the highest IoU estimate.
|
||||
If `multimask_output=False`, it's the same as `high_res_multimasks`.
|
||||
- obj_ptr: [B, C] shape, the object pointer vector for the output mask, extracted
|
||||
based on the output token from the SAM mask decoder.
|
||||
"""
|
||||
B = backbone_features.size(0)
|
||||
device = backbone_features.device
|
||||
assert backbone_features.size(1) == self.sam_prompt_embed_dim
|
||||
assert backbone_features.size(2) == self.sam_image_embedding_size
|
||||
assert backbone_features.size(3) == self.sam_image_embedding_size
|
||||
|
||||
# a) Handle point prompts
|
||||
if point_inputs is not None:
|
||||
sam_point_coords = point_inputs["point_coords"]
|
||||
sam_point_labels = point_inputs["point_labels"]
|
||||
assert sam_point_coords.size(0) == B and sam_point_labels.size(0) == B
|
||||
else:
|
||||
# If no points are provide, pad with an empty point (with label -1)
|
||||
sam_point_coords = torch.zeros(B, 1, 2, device=device)
|
||||
sam_point_labels = -torch.ones(B, 1, dtype=torch.int32, device=device)
|
||||
|
||||
# b) Handle mask prompts
|
||||
if mask_inputs is not None:
|
||||
# If mask_inputs is provided, downsize it into low-res mask input if needed
|
||||
# and feed it as a dense mask prompt into the SAM mask encoder
|
||||
assert len(mask_inputs.shape) == 4 and mask_inputs.shape[:2] == (B, 1)
|
||||
if mask_inputs.shape[-2:] != self.sam_prompt_encoder.mask_input_size:
|
||||
sam_mask_prompt = F.interpolate(
|
||||
mask_inputs.float(),
|
||||
size=self.sam_prompt_encoder.mask_input_size,
|
||||
align_corners=False,
|
||||
mode="bilinear",
|
||||
antialias=True, # use antialias for downsampling
|
||||
)
|
||||
else:
|
||||
sam_mask_prompt = mask_inputs
|
||||
else:
|
||||
# Otherwise, simply feed None (and SAM's prompt encoder will add
|
||||
# a learned `no_mask_embed` to indicate no mask input in this case).
|
||||
sam_mask_prompt = None
|
||||
|
||||
sparse_embeddings, dense_embeddings = self.sam_prompt_encoder(
|
||||
points=(sam_point_coords, sam_point_labels),
|
||||
boxes=None,
|
||||
masks=sam_mask_prompt,
|
||||
text_embeds=text_inputs
|
||||
)
|
||||
(
|
||||
low_res_multimasks,
|
||||
ious,
|
||||
sam_output_tokens,
|
||||
object_score_logits,
|
||||
) = self.sam_mask_decoder(
|
||||
image_embeddings=backbone_features,
|
||||
image_pe=self.sam_prompt_encoder.get_dense_pe(),
|
||||
sparse_prompt_embeddings=sparse_embeddings,
|
||||
dense_prompt_embeddings=dense_embeddings,
|
||||
multimask_output=multimask_output,
|
||||
repeat_image=False, # the image is already batched
|
||||
high_res_features=high_res_features,
|
||||
)
|
||||
if self.pred_obj_scores:
|
||||
is_obj_appearing = object_score_logits > 0
|
||||
|
||||
# Mask used for spatial memories is always a *hard* choice between obj and no obj,
|
||||
# consistent with the actual mask prediction
|
||||
low_res_multimasks = torch.where(
|
||||
is_obj_appearing[:, None, None],
|
||||
low_res_multimasks,
|
||||
NO_OBJ_SCORE,
|
||||
)
|
||||
|
||||
# convert masks from possibly bfloat16 (or float16) to float32
|
||||
# (older PyTorch versions before 2.1 don't support `interpolate` on bf16)
|
||||
low_res_multimasks = low_res_multimasks.float()
|
||||
high_res_multimasks = F.interpolate(
|
||||
low_res_multimasks,
|
||||
size=(self.image_size, self.image_size),
|
||||
mode="bilinear",
|
||||
align_corners=False,
|
||||
)
|
||||
|
||||
sam_output_token = sam_output_tokens[:, 0]
|
||||
if multimask_output:
|
||||
# take the best mask prediction (with the highest IoU estimation)
|
||||
best_iou_inds = torch.argmax(ious, dim=-1)
|
||||
batch_inds = torch.arange(B, device=device)
|
||||
low_res_masks = low_res_multimasks[batch_inds, best_iou_inds].unsqueeze(1)
|
||||
high_res_masks = high_res_multimasks[batch_inds, best_iou_inds].unsqueeze(1)
|
||||
if sam_output_tokens.size(1) > 1:
|
||||
sam_output_token = sam_output_tokens[batch_inds, best_iou_inds]
|
||||
else:
|
||||
low_res_masks, high_res_masks = low_res_multimasks, high_res_multimasks
|
||||
|
||||
# Extract object pointer from the SAM output token (with occlusion handling)
|
||||
obj_ptr = self.obj_ptr_proj(sam_output_token)
|
||||
if self.pred_obj_scores:
|
||||
# Allow *soft* no obj ptr, unlike for masks
|
||||
if self.soft_no_obj_ptr:
|
||||
# Only hard possible with gt
|
||||
assert not self.teacher_force_obj_scores_for_mem
|
||||
lambda_is_obj_appearing = object_score_logits.sigmoid()
|
||||
else:
|
||||
lambda_is_obj_appearing = is_obj_appearing.float()
|
||||
|
||||
if self.fixed_no_obj_ptr:
|
||||
obj_ptr = lambda_is_obj_appearing * obj_ptr
|
||||
obj_ptr = obj_ptr + (1 - lambda_is_obj_appearing) * self.no_obj_ptr
|
||||
|
||||
return (
|
||||
low_res_multimasks,
|
||||
high_res_multimasks,
|
||||
ious,
|
||||
low_res_masks,
|
||||
high_res_masks,
|
||||
obj_ptr,
|
||||
object_score_logits,
|
||||
)
|
||||
|
||||
def _use_mask_as_output(self, backbone_features, high_res_features, mask_inputs):
|
||||
"""
|
||||
Directly turn binary `mask_inputs` into a output mask logits without using SAM.
|
||||
(same input and output shapes as in _forward_sam_heads above).
|
||||
"""
|
||||
# Use -10/+10 as logits for neg/pos pixels (very close to 0/1 in prob after sigmoid).
|
||||
out_scale, out_bias = 20.0, -10.0 # sigmoid(-10.0)=4.5398e-05
|
||||
mask_inputs_float = mask_inputs.float()
|
||||
high_res_masks = mask_inputs_float * out_scale + out_bias
|
||||
low_res_masks = F.interpolate(
|
||||
high_res_masks,
|
||||
size=(high_res_masks.size(-2) // 4, high_res_masks.size(-1) // 4),
|
||||
align_corners=False,
|
||||
mode="bilinear",
|
||||
antialias=True, # use antialias for downsampling
|
||||
)
|
||||
# a dummy IoU prediction of all 1's under mask input
|
||||
ious = mask_inputs.new_ones(mask_inputs.size(0), 1).float()
|
||||
if not self.use_obj_ptrs_in_encoder:
|
||||
# all zeros as a dummy object pointer (of shape [B, C])
|
||||
obj_ptr = torch.zeros(
|
||||
mask_inputs.size(0), self.hidden_dim, device=mask_inputs.device
|
||||
)
|
||||
else:
|
||||
# produce an object pointer using the SAM decoder from the mask input
|
||||
_, _, _, _, _, obj_ptr, _ = self._forward_sam_heads(
|
||||
backbone_features=backbone_features,
|
||||
mask_inputs=self.mask_downsample(mask_inputs_float),
|
||||
high_res_features=high_res_features,
|
||||
)
|
||||
# In this method, we are treating mask_input as output, e.g. using it directly to create spatial mem;
|
||||
# Below, we follow the same design axiom to use mask_input to decide if obj appears or not instead of relying
|
||||
# on the object_scores from the SAM decoder.
|
||||
is_obj_appearing = torch.any(mask_inputs.flatten(1).float() > 0.0, dim=1)
|
||||
is_obj_appearing = is_obj_appearing[..., None]
|
||||
lambda_is_obj_appearing = is_obj_appearing.float()
|
||||
object_score_logits = out_scale * lambda_is_obj_appearing + out_bias
|
||||
if self.pred_obj_scores:
|
||||
if self.fixed_no_obj_ptr:
|
||||
obj_ptr = lambda_is_obj_appearing * obj_ptr
|
||||
obj_ptr = obj_ptr + (1 - lambda_is_obj_appearing) * self.no_obj_ptr
|
||||
|
||||
return (
|
||||
low_res_masks,
|
||||
high_res_masks,
|
||||
ious,
|
||||
low_res_masks,
|
||||
high_res_masks,
|
||||
obj_ptr,
|
||||
object_score_logits,
|
||||
)
|
||||
|
||||
def forward_image(self, img_batch: torch.Tensor):
|
||||
"""Get the image feature on the input batch."""
|
||||
backbone_out = self.image_encoder(img_batch)
|
||||
if self.use_high_res_features_in_sam:
|
||||
# precompute projected level 0 and level 1 features in SAM decoder
|
||||
# to avoid running it again on every SAM click
|
||||
backbone_out["backbone_fpn"][0] = self.sam_mask_decoder.conv_s0(
|
||||
backbone_out["backbone_fpn"][0]
|
||||
)
|
||||
backbone_out["backbone_fpn"][1] = self.sam_mask_decoder.conv_s1(
|
||||
backbone_out["backbone_fpn"][1]
|
||||
)
|
||||
return backbone_out
|
||||
|
||||
def _prepare_backbone_features(self, backbone_out):
|
||||
"""Prepare and flatten visual features."""
|
||||
backbone_out = backbone_out.copy()
|
||||
assert len(backbone_out["backbone_fpn"]) == len(backbone_out["vision_pos_enc"])
|
||||
assert len(backbone_out["backbone_fpn"]) >= self.num_feature_levels
|
||||
|
||||
feature_maps = backbone_out["backbone_fpn"][-self.num_feature_levels :]
|
||||
vision_pos_embeds = backbone_out["vision_pos_enc"][-self.num_feature_levels :]
|
||||
|
||||
feat_sizes = [(x.shape[-2], x.shape[-1]) for x in vision_pos_embeds]
|
||||
# flatten NxCxHxW to HWxNxC
|
||||
vision_feats = [x.flatten(2).permute(2, 0, 1) for x in feature_maps]
|
||||
vision_pos_embeds = [x.flatten(2).permute(2, 0, 1) for x in vision_pos_embeds]
|
||||
|
||||
return backbone_out, vision_feats, vision_pos_embeds, feat_sizes
|
||||
|
||||
def _prepare_memory_conditioned_features(
|
||||
self,
|
||||
frame_idx,
|
||||
is_init_cond_frame,
|
||||
current_vision_feats,
|
||||
current_vision_pos_embeds,
|
||||
feat_sizes,
|
||||
output_dict,
|
||||
num_frames,
|
||||
track_in_reverse=False, # tracking in reverse time order (for demo usage)
|
||||
):
|
||||
"""Fuse the current frame's visual feature map with previous memory."""
|
||||
B = current_vision_feats[-1].size(1) # batch size on this frame
|
||||
C = self.hidden_dim
|
||||
H, W = feat_sizes[-1] # top-level (lowest-resolution) feature size
|
||||
device = current_vision_feats[-1].device
|
||||
# The case of `self.num_maskmem == 0` below is primarily used for reproducing SAM on images.
|
||||
# In this case, we skip the fusion with any memory.
|
||||
if self.num_maskmem == 0: # Disable memory and skip fusion
|
||||
pix_feat = current_vision_feats[-1].permute(1, 2, 0).view(B, C, H, W)
|
||||
return pix_feat
|
||||
|
||||
num_obj_ptr_tokens = 0
|
||||
# Step 1: condition the visual features of the current frame on previous memories
|
||||
if not is_init_cond_frame:
|
||||
# Retrieve the memories encoded with the maskmem backbone
|
||||
to_cat_memory, to_cat_memory_pos_embed = [], []
|
||||
# Add conditioning frames's output first (all cond frames have t_pos=0 for
|
||||
# when getting temporal positional embedding below)
|
||||
assert len(output_dict["cond_frame_outputs"]) > 0
|
||||
# Select a maximum number of temporally closest cond frames for cross attention
|
||||
cond_outputs = output_dict["cond_frame_outputs"]
|
||||
selected_cond_outputs, unselected_cond_outputs = select_closest_cond_frames(
|
||||
frame_idx, cond_outputs, self.max_cond_frames_in_attn
|
||||
)
|
||||
t_pos_and_prevs = [(0, out) for out in selected_cond_outputs.values()]
|
||||
# Add last (self.num_maskmem - 1) frames before current frame for non-conditioning memory
|
||||
# the earliest one has t_pos=1 and the latest one has t_pos=self.num_maskmem-1
|
||||
# We also allow taking the memory frame non-consecutively (with r>1), in which case
|
||||
# we take (self.num_maskmem - 2) frames among every r-th frames plus the last frame.
|
||||
r = self.memory_temporal_stride_for_eval
|
||||
for t_pos in range(1, self.num_maskmem):
|
||||
t_rel = self.num_maskmem - t_pos # how many frames before current frame
|
||||
if t_rel == 1:
|
||||
# for t_rel == 1, we take the last frame (regardless of r)
|
||||
if not track_in_reverse:
|
||||
# the frame immediately before this frame (i.e. frame_idx - 1)
|
||||
prev_frame_idx = frame_idx - t_rel
|
||||
else:
|
||||
# the frame immediately after this frame (i.e. frame_idx + 1)
|
||||
prev_frame_idx = frame_idx + t_rel
|
||||
else:
|
||||
# for t_rel >= 2, we take the memory frame from every r-th frames
|
||||
if not track_in_reverse:
|
||||
# first find the nearest frame among every r-th frames before this frame
|
||||
# for r=1, this would be (frame_idx - 2)
|
||||
prev_frame_idx = ((frame_idx - 2) // r) * r
|
||||
# then seek further among every r-th frames
|
||||
prev_frame_idx = prev_frame_idx - (t_rel - 2) * r
|
||||
else:
|
||||
# first find the nearest frame among every r-th frames after this frame
|
||||
# for r=1, this would be (frame_idx + 2)
|
||||
prev_frame_idx = -(-(frame_idx + 2) // r) * r
|
||||
# then seek further among every r-th frames
|
||||
prev_frame_idx = prev_frame_idx + (t_rel - 2) * r
|
||||
out = output_dict["non_cond_frame_outputs"].get(prev_frame_idx, None)
|
||||
if out is None:
|
||||
# If an unselected conditioning frame is among the last (self.num_maskmem - 1)
|
||||
# frames, we still attend to it as if it's a non-conditioning frame.
|
||||
out = unselected_cond_outputs.get(prev_frame_idx, None)
|
||||
t_pos_and_prevs.append((t_pos, out))
|
||||
|
||||
for t_pos, prev in t_pos_and_prevs:
|
||||
if prev is None:
|
||||
continue # skip padding frames
|
||||
# "maskmem_features" might have been offloaded to CPU in demo use cases,
|
||||
# so we load it back to GPU (it's a no-op if it's already on GPU).
|
||||
feats = prev["maskmem_features"].cuda(non_blocking=True)
|
||||
to_cat_memory.append(feats.flatten(2).permute(2, 0, 1))
|
||||
# Spatial positional encoding (it might have been offloaded to CPU in eval)
|
||||
maskmem_enc = prev["maskmem_pos_enc"][-1].cuda()
|
||||
maskmem_enc = maskmem_enc.flatten(2).permute(2, 0, 1)
|
||||
# Temporal positional encoding
|
||||
maskmem_enc = (
|
||||
maskmem_enc + self.maskmem_tpos_enc[self.num_maskmem - t_pos - 1]
|
||||
)
|
||||
to_cat_memory_pos_embed.append(maskmem_enc)
|
||||
|
||||
# Construct the list of past object pointers
|
||||
if self.use_obj_ptrs_in_encoder:
|
||||
max_obj_ptrs_in_encoder = min(num_frames, self.max_obj_ptrs_in_encoder)
|
||||
# First add those object pointers from selected conditioning frames
|
||||
# (optionally, only include object pointers in the past during evaluation)
|
||||
if not self.training and self.only_obj_ptrs_in_the_past_for_eval:
|
||||
ptr_cond_outputs = {
|
||||
t: out
|
||||
for t, out in selected_cond_outputs.items()
|
||||
if (t >= frame_idx if track_in_reverse else t <= frame_idx)
|
||||
}
|
||||
else:
|
||||
ptr_cond_outputs = selected_cond_outputs
|
||||
pos_and_ptrs = [
|
||||
# Temporal pos encoding contains how far away each pointer is from current frame
|
||||
(abs(frame_idx - t), out["obj_ptr"])
|
||||
for t, out in ptr_cond_outputs.items()
|
||||
]
|
||||
# Add up to (max_obj_ptrs_in_encoder - 1) non-conditioning frames before current frame
|
||||
for t_diff in range(1, max_obj_ptrs_in_encoder):
|
||||
t = frame_idx + t_diff if track_in_reverse else frame_idx - t_diff
|
||||
if t < 0 or (num_frames is not None and t >= num_frames):
|
||||
break
|
||||
out = output_dict["non_cond_frame_outputs"].get(
|
||||
t, unselected_cond_outputs.get(t, None)
|
||||
)
|
||||
if out is not None:
|
||||
pos_and_ptrs.append((t_diff, out["obj_ptr"]))
|
||||
# If we have at least one object pointer, add them to the across attention
|
||||
if len(pos_and_ptrs) > 0:
|
||||
pos_list, ptrs_list = zip(*pos_and_ptrs)
|
||||
# stack object pointers along dim=0 into [ptr_seq_len, B, C] shape
|
||||
obj_ptrs = torch.stack(ptrs_list, dim=0)
|
||||
# a temporal positional embedding based on how far each object pointer is from
|
||||
# the current frame (sine embedding normalized by the max pointer num).
|
||||
if self.add_tpos_enc_to_obj_ptrs:
|
||||
t_diff_max = max_obj_ptrs_in_encoder - 1
|
||||
tpos_dim = C if self.proj_tpos_enc_in_obj_ptrs else self.mem_dim
|
||||
obj_pos = torch.tensor(pos_list, device=device)
|
||||
obj_pos = get_1d_sine_pe(obj_pos / t_diff_max, dim=tpos_dim)
|
||||
obj_pos = self.obj_ptr_tpos_proj(obj_pos)
|
||||
obj_pos = obj_pos.unsqueeze(1).expand(-1, B, self.mem_dim)
|
||||
else:
|
||||
obj_pos = obj_ptrs.new_zeros(len(pos_list), B, self.mem_dim)
|
||||
if self.mem_dim < C:
|
||||
# split a pointer into (C // self.mem_dim) tokens for self.mem_dim < C
|
||||
obj_ptrs = obj_ptrs.reshape(
|
||||
-1, B, C // self.mem_dim, self.mem_dim
|
||||
)
|
||||
obj_ptrs = obj_ptrs.permute(0, 2, 1, 3).flatten(0, 1)
|
||||
obj_pos = obj_pos.repeat_interleave(C // self.mem_dim, dim=0)
|
||||
to_cat_memory.append(obj_ptrs)
|
||||
to_cat_memory_pos_embed.append(obj_pos)
|
||||
num_obj_ptr_tokens = obj_ptrs.shape[0]
|
||||
else:
|
||||
num_obj_ptr_tokens = 0
|
||||
else:
|
||||
# for initial conditioning frames, encode them without using any previous memory
|
||||
if self.directly_add_no_mem_embed:
|
||||
# directly add no-mem embedding (instead of using the transformer encoder)
|
||||
pix_feat_with_mem = current_vision_feats[-1] + self.no_mem_embed
|
||||
pix_feat_with_mem = pix_feat_with_mem.permute(1, 2, 0).view(B, C, H, W)
|
||||
return pix_feat_with_mem
|
||||
|
||||
# Use a dummy token on the first frame (to avoid emtpy memory input to tranformer encoder)
|
||||
to_cat_memory = [self.no_mem_embed.expand(1, B, self.mem_dim)]
|
||||
to_cat_memory_pos_embed = [self.no_mem_pos_enc.expand(1, B, self.mem_dim)]
|
||||
|
||||
# Step 2: Concatenate the memories and forward through the transformer encoder
|
||||
memory = torch.cat(to_cat_memory, dim=0)
|
||||
memory_pos_embed = torch.cat(to_cat_memory_pos_embed, dim=0)
|
||||
|
||||
pix_feat_with_mem = self.memory_attention(
|
||||
curr=current_vision_feats,
|
||||
curr_pos=current_vision_pos_embeds,
|
||||
memory=memory,
|
||||
memory_pos=memory_pos_embed,
|
||||
num_obj_ptr_tokens=num_obj_ptr_tokens,
|
||||
)
|
||||
# reshape the output (HW)BC => BCHW
|
||||
pix_feat_with_mem = pix_feat_with_mem.permute(1, 2, 0).view(B, C, H, W)
|
||||
return pix_feat_with_mem
|
||||
|
||||
def _encode_new_memory(
|
||||
self,
|
||||
current_vision_feats,
|
||||
feat_sizes,
|
||||
pred_masks_high_res,
|
||||
is_mask_from_pts,
|
||||
):
|
||||
"""Encode the current image and its prediction into a memory feature."""
|
||||
B = current_vision_feats[-1].size(1) # batch size on this frame
|
||||
C = self.hidden_dim
|
||||
H, W = feat_sizes[-1] # top-level (lowest-resolution) feature size
|
||||
# top-level feature, (HW)BC => BCHW
|
||||
pix_feat = current_vision_feats[-1].permute(1, 2, 0).view(B, C, H, W)
|
||||
if self.non_overlap_masks_for_mem_enc and not self.training:
|
||||
# optionally, apply non-overlapping constraints to the masks (it's applied
|
||||
# in the batch dimension and should only be used during eval, where all
|
||||
# the objects come from the same video under batch size 1).
|
||||
pred_masks_high_res = self._apply_non_overlapping_constraints(
|
||||
pred_masks_high_res
|
||||
)
|
||||
# scale the raw mask logits with a temperature before applying sigmoid
|
||||
binarize = self.binarize_mask_from_pts_for_mem_enc and is_mask_from_pts
|
||||
if binarize and not self.training:
|
||||
mask_for_mem = (pred_masks_high_res > 0).float()
|
||||
else:
|
||||
# apply sigmoid on the raw mask logits to turn them into range (0, 1)
|
||||
mask_for_mem = torch.sigmoid(pred_masks_high_res)
|
||||
# apply scale and bias terms to the sigmoid probabilities
|
||||
if self.sigmoid_scale_for_mem_enc != 1.0:
|
||||
mask_for_mem = mask_for_mem * self.sigmoid_scale_for_mem_enc
|
||||
if self.sigmoid_bias_for_mem_enc != 0.0:
|
||||
mask_for_mem = mask_for_mem + self.sigmoid_bias_for_mem_enc
|
||||
maskmem_out = self.memory_encoder(
|
||||
pix_feat, mask_for_mem, skip_mask_sigmoid=True # sigmoid already applied
|
||||
)
|
||||
maskmem_features = maskmem_out["vision_features"]
|
||||
maskmem_pos_enc = maskmem_out["vision_pos_enc"]
|
||||
|
||||
return maskmem_features, maskmem_pos_enc
|
||||
|
||||
def track_step(
|
||||
self,
|
||||
frame_idx,
|
||||
is_init_cond_frame,
|
||||
current_vision_feats,
|
||||
current_vision_pos_embeds,
|
||||
feat_sizes,
|
||||
point_inputs,
|
||||
mask_inputs,
|
||||
output_dict,
|
||||
num_frames,
|
||||
track_in_reverse=False, # tracking in reverse time order (for demo usage)
|
||||
# Whether to run the memory encoder on the predicted masks. Sometimes we might want
|
||||
# to skip the memory encoder with `run_mem_encoder=False`. For example,
|
||||
# in demo we might call `track_step` multiple times for each user click,
|
||||
# and only encode the memory when the user finalizes their clicks. And in ablation
|
||||
# settings like SAM training on static images, we don't need the memory encoder.
|
||||
run_mem_encoder=True,
|
||||
# The previously predicted SAM mask logits (which can be fed together with new clicks in demo).
|
||||
prev_sam_mask_logits=None,
|
||||
text_inputs=None,
|
||||
):
|
||||
current_out = {"point_inputs": point_inputs, "mask_inputs": mask_inputs}
|
||||
# High-resolution feature maps for the SAM head, reshape (HW)BC => BCHW
|
||||
if len(current_vision_feats) > 1:
|
||||
high_res_features = [
|
||||
x.permute(1, 2, 0).view(x.size(1), x.size(2), *s)
|
||||
for x, s in zip(current_vision_feats[:-1], feat_sizes[:-1])
|
||||
]
|
||||
else:
|
||||
high_res_features = None
|
||||
if mask_inputs is not None and self.use_mask_input_as_output_without_sam:
|
||||
# When use_mask_input_as_output_without_sam=True, we directly output the mask input
|
||||
# (see it as a GT mask) without using a SAM prompt encoder + mask decoder.
|
||||
pix_feat = current_vision_feats[-1].permute(1, 2, 0)
|
||||
pix_feat = pix_feat.view(-1, self.hidden_dim, *feat_sizes[-1])
|
||||
sam_outputs = self._use_mask_as_output(
|
||||
pix_feat, high_res_features, mask_inputs
|
||||
)
|
||||
else:
|
||||
# fused the visual feature with previous memory features in the memory bank
|
||||
pix_feat_with_mem = self._prepare_memory_conditioned_features(
|
||||
frame_idx=frame_idx,
|
||||
is_init_cond_frame=is_init_cond_frame,
|
||||
current_vision_feats=current_vision_feats[-1:],
|
||||
current_vision_pos_embeds=current_vision_pos_embeds[-1:],
|
||||
feat_sizes=feat_sizes[-1:],
|
||||
output_dict=output_dict,
|
||||
num_frames=num_frames,
|
||||
track_in_reverse=track_in_reverse,
|
||||
)
|
||||
# apply SAM-style segmentation head
|
||||
# here we might feed previously predicted low-res SAM mask logits into the SAM mask decoder,
|
||||
# e.g. in demo where such logits come from earlier interaction instead of correction sampling
|
||||
# (in this case, any `mask_inputs` shouldn't reach here as they are sent to _use_mask_as_output instead)
|
||||
if prev_sam_mask_logits is not None:
|
||||
assert point_inputs is not None and mask_inputs is None
|
||||
mask_inputs = prev_sam_mask_logits
|
||||
multimask_output = self._use_multimask(is_init_cond_frame, point_inputs)
|
||||
sam_outputs = self._forward_sam_heads(
|
||||
backbone_features=pix_feat_with_mem,
|
||||
point_inputs=point_inputs,
|
||||
mask_inputs=mask_inputs,
|
||||
high_res_features=high_res_features,
|
||||
multimask_output=multimask_output,
|
||||
text_inputs=text_inputs
|
||||
)
|
||||
(
|
||||
_,
|
||||
_,
|
||||
_,
|
||||
low_res_masks,
|
||||
high_res_masks,
|
||||
obj_ptr,
|
||||
_,
|
||||
) = sam_outputs
|
||||
|
||||
current_out["pred_masks"] = low_res_masks
|
||||
current_out["pred_masks_high_res"] = high_res_masks
|
||||
current_out["obj_ptr"] = obj_ptr
|
||||
|
||||
# Finally run the memory encoder on the predicted mask to encode
|
||||
# it into a new memory feature (that can be used in future frames)
|
||||
if run_mem_encoder and self.num_maskmem > 0:
|
||||
high_res_masks_for_mem_enc = high_res_masks
|
||||
maskmem_features, maskmem_pos_enc = self._encode_new_memory(
|
||||
current_vision_feats=current_vision_feats,
|
||||
feat_sizes=feat_sizes,
|
||||
pred_masks_high_res=high_res_masks_for_mem_enc,
|
||||
is_mask_from_pts=(point_inputs is not None),
|
||||
)
|
||||
current_out["maskmem_features"] = maskmem_features
|
||||
current_out["maskmem_pos_enc"] = maskmem_pos_enc
|
||||
else:
|
||||
current_out["maskmem_features"] = None
|
||||
current_out["maskmem_pos_enc"] = None
|
||||
|
||||
return current_out
|
||||
|
||||
def _use_multimask(self, is_init_cond_frame, point_inputs):
|
||||
"""Whether to use multimask output in the SAM head."""
|
||||
num_pts = 0 if point_inputs is None else point_inputs["point_labels"].size(1)
|
||||
multimask_output = (
|
||||
self.multimask_output_in_sam
|
||||
and (is_init_cond_frame or self.multimask_output_for_tracking)
|
||||
and (self.multimask_min_pt_num <= num_pts <= self.multimask_max_pt_num)
|
||||
)
|
||||
return multimask_output
|
||||
|
||||
def _apply_non_overlapping_constraints(self, pred_masks):
|
||||
"""
|
||||
Apply non-overlapping constraints to the object scores in pred_masks. Here we
|
||||
keep only the highest scoring object at each spatial location in pred_masks.
|
||||
"""
|
||||
batch_size = pred_masks.size(0)
|
||||
if batch_size == 1:
|
||||
return pred_masks
|
||||
|
||||
device = pred_masks.device
|
||||
# "max_obj_inds": object index of the object with the highest score at each location
|
||||
max_obj_inds = torch.argmax(pred_masks, dim=0, keepdim=True)
|
||||
# "batch_obj_inds": object index of each object slice (along dim 0) in `pred_masks`
|
||||
batch_obj_inds = torch.arange(batch_size, device=device)[:, None, None, None]
|
||||
keep = max_obj_inds == batch_obj_inds
|
||||
# suppress overlapping regions' scores below -10.0 so that the foreground regions
|
||||
# don't overlap (here sigmoid(-10.0)=4.5398e-05)
|
||||
pred_masks = torch.where(keep, pred_masks, torch.clamp(pred_masks, max=-10.0))
|
||||
return pred_masks
|
||||
@@ -0,0 +1,149 @@
|
||||
# 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 copy
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
def select_closest_cond_frames(frame_idx, cond_frame_outputs, max_cond_frame_num):
|
||||
"""
|
||||
Select up to `max_cond_frame_num` conditioning frames from `cond_frame_outputs`
|
||||
that are temporally closest to the current frame at `frame_idx`. Here, we take
|
||||
- a) the closest conditioning frame before `frame_idx` (if any);
|
||||
- b) the closest conditioning frame after `frame_idx` (if any);
|
||||
- c) any other temporally closest conditioning frames until reaching a total
|
||||
of `max_cond_frame_num` conditioning frames.
|
||||
|
||||
Outputs:
|
||||
- selected_outputs: selected items (keys & values) from `cond_frame_outputs`.
|
||||
- unselected_outputs: items (keys & values) not selected in `cond_frame_outputs`.
|
||||
"""
|
||||
if max_cond_frame_num == -1 or len(cond_frame_outputs) <= max_cond_frame_num:
|
||||
selected_outputs = cond_frame_outputs
|
||||
unselected_outputs = {}
|
||||
else:
|
||||
assert max_cond_frame_num >= 2, "we should allow using 2+ conditioning frames"
|
||||
selected_outputs = {}
|
||||
|
||||
# the closest conditioning frame before `frame_idx` (if any)
|
||||
idx_before = max((t for t in cond_frame_outputs if t < frame_idx), default=None)
|
||||
if idx_before is not None:
|
||||
selected_outputs[idx_before] = cond_frame_outputs[idx_before]
|
||||
|
||||
# the closest conditioning frame after `frame_idx` (if any)
|
||||
idx_after = min((t for t in cond_frame_outputs if t >= frame_idx), default=None)
|
||||
if idx_after is not None:
|
||||
selected_outputs[idx_after] = cond_frame_outputs[idx_after]
|
||||
|
||||
# add other temporally closest conditioning frames until reaching a total
|
||||
# of `max_cond_frame_num` conditioning frames.
|
||||
num_remain = max_cond_frame_num - len(selected_outputs)
|
||||
inds_remain = sorted(
|
||||
(t for t in cond_frame_outputs if t not in selected_outputs),
|
||||
key=lambda x: abs(x - frame_idx),
|
||||
)[:num_remain]
|
||||
selected_outputs.update((t, cond_frame_outputs[t]) for t in inds_remain)
|
||||
unselected_outputs = {
|
||||
t: v for t, v in cond_frame_outputs.items() if t not in selected_outputs
|
||||
}
|
||||
|
||||
return selected_outputs, unselected_outputs
|
||||
|
||||
|
||||
def get_1d_sine_pe(pos_inds, dim, temperature=10000):
|
||||
"""
|
||||
Get 1D sine positional embedding as in the original Transformer paper.
|
||||
"""
|
||||
pe_dim = dim // 2
|
||||
dim_t = torch.arange(pe_dim, dtype=torch.float32, device=pos_inds.device)
|
||||
dim_t = temperature ** (2 * (dim_t // 2) / pe_dim)
|
||||
|
||||
pos_embed = pos_inds.unsqueeze(-1) / dim_t
|
||||
pos_embed = torch.cat([pos_embed.sin(), pos_embed.cos()], dim=-1)
|
||||
return pos_embed
|
||||
|
||||
|
||||
def get_activation_fn(activation):
|
||||
"""Return an activation function given a string"""
|
||||
if activation == "relu":
|
||||
return F.relu
|
||||
if activation == "gelu":
|
||||
return F.gelu
|
||||
if activation == "glu":
|
||||
return F.glu
|
||||
raise RuntimeError(f"activation should be relu/gelu, not {activation}.")
|
||||
|
||||
|
||||
def get_clones(module, N):
|
||||
return nn.ModuleList([copy.deepcopy(module) for i in range(N)])
|
||||
|
||||
|
||||
class DropPath(nn.Module):
|
||||
# adapted from https://github.com/huggingface/pytorch-image-models/blob/main/timm/layers/drop.py
|
||||
def __init__(self, drop_prob=0.0, scale_by_keep=True):
|
||||
super(DropPath, self).__init__()
|
||||
self.drop_prob = drop_prob
|
||||
self.scale_by_keep = scale_by_keep
|
||||
|
||||
def forward(self, x):
|
||||
if self.drop_prob == 0.0 or not self.training:
|
||||
return x
|
||||
keep_prob = 1 - self.drop_prob
|
||||
shape = (x.shape[0],) + (1,) * (x.ndim - 1)
|
||||
random_tensor = x.new_empty(shape).bernoulli_(keep_prob)
|
||||
if keep_prob > 0.0 and self.scale_by_keep:
|
||||
random_tensor.div_(keep_prob)
|
||||
return x * random_tensor
|
||||
|
||||
|
||||
# Lightly adapted from
|
||||
# https://github.com/facebookresearch/MaskFormer/blob/main/mask_former/modeling/transformer/transformer_predictor.py # noqa
|
||||
class MLP(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_dim: int,
|
||||
hidden_dim: int,
|
||||
output_dim: int,
|
||||
num_layers: int,
|
||||
activation: nn.Module = nn.ReLU,
|
||||
sigmoid_output: bool = False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.num_layers = num_layers
|
||||
h = [hidden_dim] * (num_layers - 1)
|
||||
self.layers = nn.ModuleList(
|
||||
nn.Linear(n, k) for n, k in zip([input_dim] + h, h + [output_dim])
|
||||
)
|
||||
self.sigmoid_output = sigmoid_output
|
||||
self.act = activation()
|
||||
|
||||
def forward(self, x):
|
||||
for i, layer in enumerate(self.layers):
|
||||
x = self.act(layer(x)) if i < self.num_layers - 1 else layer(x)
|
||||
if self.sigmoid_output:
|
||||
x = F.sigmoid(x)
|
||||
return x
|
||||
|
||||
|
||||
# From https://github.com/facebookresearch/detectron2/blob/main/detectron2/layers/batch_norm.py # noqa
|
||||
# Itself from https://github.com/facebookresearch/ConvNeXt/blob/d1fa8f6fef0a165b27399986cc2bdacc92777e40/models/convnext.py#L119 # noqa
|
||||
class LayerNorm2d(nn.Module):
|
||||
def __init__(self, num_channels: int, eps: float = 1e-6) -> None:
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.ones(num_channels))
|
||||
self.bias = nn.Parameter(torch.zeros(num_channels))
|
||||
self.eps = eps
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
u = x.mean(1, keepdim=True)
|
||||
s = (x - u).pow(2).mean(1, keepdim=True)
|
||||
x = (x - u) / torch.sqrt(s + self.eps)
|
||||
x = self.weight[:, None, None] * x + self.bias[:, None, None]
|
||||
return x
|
||||
@@ -0,0 +1,446 @@
|
||||
# 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 logging
|
||||
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL.Image import Image
|
||||
|
||||
from model.segment_anything_2.sam2.modeling.sam2_base import SAM2Base
|
||||
|
||||
from model.segment_anything_2.sam2.utils.transforms import SAM2Transforms
|
||||
|
||||
|
||||
class SAM2ImagePredictor:
|
||||
def __init__(
|
||||
self,
|
||||
sam_model: SAM2Base,
|
||||
mask_threshold=0.0,
|
||||
max_hole_area=0.0,
|
||||
max_sprinkle_area=0.0,
|
||||
) -> None:
|
||||
"""
|
||||
Uses SAM-2 to calculate the image embedding for an image, and then
|
||||
allow repeated, efficient mask prediction given prompts.
|
||||
|
||||
Arguments:
|
||||
sam_model (Sam-2): The model to use for mask prediction.
|
||||
mask_threshold (float): The threshold to use when converting mask logits
|
||||
to binary masks. Masks are thresholded at 0 by default.
|
||||
fill_hole_area (int): If fill_hole_area > 0, we fill small holes in up to
|
||||
the maximum area of fill_hole_area in low_res_masks.
|
||||
"""
|
||||
super().__init__()
|
||||
self.model = sam_model
|
||||
self._transforms = SAM2Transforms(
|
||||
resolution=self.model.image_size,
|
||||
mask_threshold=mask_threshold,
|
||||
max_hole_area=max_hole_area,
|
||||
max_sprinkle_area=max_sprinkle_area,
|
||||
)
|
||||
|
||||
# Predictor state
|
||||
self._is_image_set = False
|
||||
self._features = None
|
||||
self._orig_hw = None
|
||||
# Whether the predictor is set for single image or a batch of images
|
||||
self._is_batch = False
|
||||
|
||||
# Predictor config
|
||||
self.mask_threshold = mask_threshold
|
||||
|
||||
# Spatial dim for backbone feature maps
|
||||
self._bb_feat_sizes = [
|
||||
(256, 256),
|
||||
(128, 128),
|
||||
(64, 64),
|
||||
]
|
||||
|
||||
@torch.no_grad()
|
||||
def set_image(
|
||||
self,
|
||||
image: Union[np.ndarray, Image],
|
||||
) -> None:
|
||||
"""
|
||||
Calculates the image embeddings for the provided image, allowing
|
||||
masks to be predicted with the 'predict' method.
|
||||
|
||||
Arguments:
|
||||
image (np.ndarray or PIL Image): The input image to embed in RGB format. The image should be in HWC format if np.ndarray, or WHC format if PIL Image
|
||||
with pixel values in [0, 255].
|
||||
image_format (str): The color format of the image, in ['RGB', 'BGR'].
|
||||
"""
|
||||
self.reset_predictor()
|
||||
# Transform the image to the form expected by the model
|
||||
if isinstance(image, np.ndarray):
|
||||
logging.info("For numpy array image, we assume (HxWxC) format")
|
||||
self._orig_hw = [image.shape[:2]]
|
||||
elif isinstance(image, Image):
|
||||
w, h = image.size
|
||||
self._orig_hw = [(h, w)]
|
||||
else:
|
||||
raise NotImplementedError("Image format not supported")
|
||||
|
||||
input_image = self._transforms(image)
|
||||
input_image = input_image[None, ...].to(self.device)
|
||||
|
||||
assert (
|
||||
len(input_image.shape) == 4 and input_image.shape[1] == 3
|
||||
), f"input_image must be of size 1x3xHxW, got {input_image.shape}"
|
||||
logging.info("Computing image embeddings for the provided image...")
|
||||
backbone_out = self.model.forward_image(input_image)
|
||||
_, vision_feats, _, _ = self.model._prepare_backbone_features(backbone_out)
|
||||
# Add no_mem_embed, which is added to the lowest rest feat. map during training on videos
|
||||
if self.model.directly_add_no_mem_embed:
|
||||
vision_feats[-1] = vision_feats[-1] + self.model.no_mem_embed
|
||||
|
||||
feats = [
|
||||
feat.permute(1, 2, 0).view(1, -1, *feat_size)
|
||||
for feat, feat_size in zip(vision_feats[::-1], self._bb_feat_sizes[::-1])
|
||||
][::-1]
|
||||
self._features = {"image_embed": feats[-1], "high_res_feats": feats[:-1]}
|
||||
self._is_image_set = True
|
||||
logging.info("Image embeddings computed.")
|
||||
|
||||
@torch.no_grad()
|
||||
def set_image_batch(
|
||||
self,
|
||||
image_list: List[Union[np.ndarray]],
|
||||
) -> None:
|
||||
"""
|
||||
Calculates the image embeddings for the provided image batch, allowing
|
||||
masks to be predicted with the 'predict_batch' method.
|
||||
|
||||
Arguments:
|
||||
image_list (List[np.ndarray]): The input images to embed in RGB format. The image should be in HWC format if np.ndarray
|
||||
with pixel values in [0, 255].
|
||||
"""
|
||||
self.reset_predictor()
|
||||
assert isinstance(image_list, list)
|
||||
self._orig_hw = []
|
||||
for image in image_list:
|
||||
assert isinstance(
|
||||
image, np.ndarray
|
||||
), "Images are expected to be an np.ndarray in RGB format, and of shape HWC"
|
||||
self._orig_hw.append(image.shape[:2])
|
||||
# Transform the image to the form expected by the model
|
||||
img_batch = self._transforms.forward_batch(image_list)
|
||||
img_batch = img_batch.to(self.device)
|
||||
batch_size = img_batch.shape[0]
|
||||
assert (
|
||||
len(img_batch.shape) == 4 and img_batch.shape[1] == 3
|
||||
), f"img_batch must be of size Bx3xHxW, got {img_batch.shape}"
|
||||
logging.info("Computing image embeddings for the provided images...")
|
||||
backbone_out = self.model.forward_image(img_batch)
|
||||
_, vision_feats, _, _ = self.model._prepare_backbone_features(backbone_out)
|
||||
# Add no_mem_embed, which is added to the lowest rest feat. map during training on videos
|
||||
if self.model.directly_add_no_mem_embed:
|
||||
vision_feats[-1] = vision_feats[-1] + self.model.no_mem_embed
|
||||
|
||||
feats = [
|
||||
feat.permute(1, 2, 0).view(batch_size, -1, *feat_size)
|
||||
for feat, feat_size in zip(vision_feats[::-1], self._bb_feat_sizes[::-1])
|
||||
][::-1]
|
||||
self._features = {"image_embed": feats[-1], "high_res_feats": feats[:-1]}
|
||||
self._is_image_set = True
|
||||
self._is_batch = True
|
||||
logging.info("Image embeddings computed.")
|
||||
|
||||
def predict_batch(
|
||||
self,
|
||||
point_coords_batch: List[np.ndarray] = None,
|
||||
point_labels_batch: List[np.ndarray] = None,
|
||||
box_batch: List[np.ndarray] = None,
|
||||
mask_input_batch: List[np.ndarray] = None,
|
||||
multimask_output: bool = True,
|
||||
return_logits: bool = False,
|
||||
normalize_coords=True,
|
||||
) -> Tuple[List[np.ndarray], List[np.ndarray], List[np.ndarray]]:
|
||||
"""This function is very similar to predict(...), however it is used for batched mode, when the model is expected to generate predictions on multiple images.
|
||||
It returns a tupele of lists of masks, ious, and low_res_masks_logits.
|
||||
"""
|
||||
assert self._is_batch, "This function should only be used when in batched mode"
|
||||
if not self._is_image_set:
|
||||
raise RuntimeError(
|
||||
"An image must be set with .set_image_batch(...) before mask prediction."
|
||||
)
|
||||
num_images = len(self._features["image_embed"])
|
||||
all_masks = []
|
||||
all_ious = []
|
||||
all_low_res_masks = []
|
||||
for img_idx in range(num_images):
|
||||
# Transform input prompts
|
||||
point_coords = (
|
||||
point_coords_batch[img_idx] if point_coords_batch is not None else None
|
||||
)
|
||||
point_labels = (
|
||||
point_labels_batch[img_idx] if point_labels_batch is not None else None
|
||||
)
|
||||
box = box_batch[img_idx] if box_batch is not None else None
|
||||
mask_input = (
|
||||
mask_input_batch[img_idx] if mask_input_batch is not None else None
|
||||
)
|
||||
mask_input, unnorm_coords, labels, unnorm_box = self._prep_prompts(
|
||||
point_coords,
|
||||
point_labels,
|
||||
box,
|
||||
mask_input,
|
||||
normalize_coords,
|
||||
img_idx=img_idx,
|
||||
)
|
||||
masks, iou_predictions, low_res_masks = self._predict(
|
||||
unnorm_coords,
|
||||
labels,
|
||||
unnorm_box,
|
||||
mask_input,
|
||||
multimask_output,
|
||||
return_logits=return_logits,
|
||||
img_idx=img_idx,
|
||||
)
|
||||
masks_np = masks.squeeze(0).float().detach().cpu().numpy()
|
||||
iou_predictions_np = (
|
||||
iou_predictions.squeeze(0).float().detach().cpu().numpy()
|
||||
)
|
||||
low_res_masks_np = low_res_masks.squeeze(0).float().detach().cpu().numpy()
|
||||
all_masks.append(masks_np)
|
||||
all_ious.append(iou_predictions_np)
|
||||
all_low_res_masks.append(low_res_masks_np)
|
||||
|
||||
return all_masks, all_ious, all_low_res_masks
|
||||
|
||||
def predict(
|
||||
self,
|
||||
point_coords: Optional[np.ndarray] = None,
|
||||
point_labels: Optional[np.ndarray] = None,
|
||||
box: Optional[np.ndarray] = None,
|
||||
mask_input: Optional[np.ndarray] = None,
|
||||
multimask_output: bool = True,
|
||||
return_logits: bool = False,
|
||||
normalize_coords=True,
|
||||
) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
|
||||
"""
|
||||
Predict masks for the given input prompts, using the currently set image.
|
||||
|
||||
Arguments:
|
||||
point_coords (np.ndarray or None): A Nx2 array of point prompts to the
|
||||
model. Each point is in (X,Y) in pixels.
|
||||
point_labels (np.ndarray or None): A length N array of labels for the
|
||||
point prompts. 1 indicates a foreground point and 0 indicates a
|
||||
background point.
|
||||
box (np.ndarray or None): A length 4 array given a box prompt to the
|
||||
model, in XYXY format.
|
||||
mask_input (np.ndarray): A low resolution mask input to the model, typically
|
||||
coming from a previous prediction iteration. Has form 1xHxW, where
|
||||
for SAM, H=W=256.
|
||||
multimask_output (bool): If true, the model will return three masks.
|
||||
For ambiguous input prompts (such as a single click), this will often
|
||||
produce better masks than a single prediction. If only a single
|
||||
mask is needed, the model's predicted quality score can be used
|
||||
to select the best mask. For non-ambiguous prompts, such as multiple
|
||||
input prompts, multimask_output=False can give better results.
|
||||
return_logits (bool): If true, returns un-thresholded masks logits
|
||||
instead of a binary mask.
|
||||
normalize_coords (bool): If true, the point coordinates will be normalized to the range [0,1] and point_coords is expected to be wrt. image dimensions.
|
||||
|
||||
Returns:
|
||||
(np.ndarray): The output masks in CxHxW format, where C is the
|
||||
number of masks, and (H, W) is the original image size.
|
||||
(np.ndarray): An array of length C containing the model's
|
||||
predictions for the quality of each mask.
|
||||
(np.ndarray): An array of shape CxHxW, where C is the number
|
||||
of masks and H=W=256. These low resolution logits can be passed to
|
||||
a subsequent iteration as mask input.
|
||||
"""
|
||||
if not self._is_image_set:
|
||||
raise RuntimeError(
|
||||
"An image must be set with .set_image(...) before mask prediction."
|
||||
)
|
||||
|
||||
# Transform input prompts
|
||||
|
||||
mask_input, unnorm_coords, labels, unnorm_box = self._prep_prompts(
|
||||
point_coords, point_labels, box, mask_input, normalize_coords
|
||||
)
|
||||
|
||||
masks, iou_predictions, low_res_masks = self._predict(
|
||||
unnorm_coords,
|
||||
labels,
|
||||
unnorm_box,
|
||||
mask_input,
|
||||
multimask_output,
|
||||
return_logits=return_logits,
|
||||
)
|
||||
|
||||
masks_np = masks.squeeze(0).float().detach().cpu().numpy()
|
||||
iou_predictions_np = iou_predictions.squeeze(0).float().detach().cpu().numpy()
|
||||
low_res_masks_np = low_res_masks.squeeze(0).float().detach().cpu().numpy()
|
||||
return masks_np, iou_predictions_np, low_res_masks_np
|
||||
|
||||
def _prep_prompts(
|
||||
self, point_coords, point_labels, box, mask_logits, normalize_coords, img_idx=-1
|
||||
):
|
||||
|
||||
unnorm_coords, labels, unnorm_box, mask_input = None, None, None, None
|
||||
if point_coords is not None:
|
||||
assert (
|
||||
point_labels is not None
|
||||
), "point_labels must be supplied if point_coords is supplied."
|
||||
point_coords = torch.as_tensor(
|
||||
point_coords, dtype=torch.float, device=self.device
|
||||
)
|
||||
unnorm_coords = self._transforms.transform_coords(
|
||||
point_coords, normalize=normalize_coords, orig_hw=self._orig_hw[img_idx]
|
||||
)
|
||||
labels = torch.as_tensor(point_labels, dtype=torch.int, device=self.device)
|
||||
if len(unnorm_coords.shape) == 2:
|
||||
unnorm_coords, labels = unnorm_coords[None, ...], labels[None, ...]
|
||||
if box is not None:
|
||||
box = torch.as_tensor(box, dtype=torch.float, device=self.device)
|
||||
unnorm_box = self._transforms.transform_boxes(
|
||||
box, normalize=normalize_coords, orig_hw=self._orig_hw[img_idx]
|
||||
) # Bx2x2
|
||||
if mask_logits is not None:
|
||||
mask_input = torch.as_tensor(
|
||||
mask_logits, dtype=torch.float, device=self.device
|
||||
)
|
||||
if len(mask_input.shape) == 3:
|
||||
mask_input = mask_input[None, :, :, :]
|
||||
return mask_input, unnorm_coords, labels, unnorm_box
|
||||
|
||||
@torch.no_grad()
|
||||
def _predict(
|
||||
self,
|
||||
point_coords: Optional[torch.Tensor],
|
||||
point_labels: Optional[torch.Tensor],
|
||||
boxes: Optional[torch.Tensor] = None,
|
||||
mask_input: Optional[torch.Tensor] = None,
|
||||
multimask_output: bool = True,
|
||||
return_logits: bool = False,
|
||||
img_idx: int = -1,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Predict masks for the given input prompts, using the currently set image.
|
||||
Input prompts are batched torch tensors and are expected to already be
|
||||
transformed to the input frame using SAM2Transforms.
|
||||
|
||||
Arguments:
|
||||
point_coords (torch.Tensor or None): A BxNx2 array of point prompts to the
|
||||
model. Each point is in (X,Y) in pixels.
|
||||
point_labels (torch.Tensor or None): A BxN array of labels for the
|
||||
point prompts. 1 indicates a foreground point and 0 indicates a
|
||||
background point.
|
||||
boxes (np.ndarray or None): A Bx4 array given a box prompt to the
|
||||
model, in XYXY format.
|
||||
mask_input (np.ndarray): A low resolution mask input to the model, typically
|
||||
coming from a previous prediction iteration. Has form Bx1xHxW, where
|
||||
for SAM, H=W=256. Masks returned by a previous iteration of the
|
||||
predict method do not need further transformation.
|
||||
multimask_output (bool): If true, the model will return three masks.
|
||||
For ambiguous input prompts (such as a single click), this will often
|
||||
produce better masks than a single prediction. If only a single
|
||||
mask is needed, the model's predicted quality score can be used
|
||||
to select the best mask. For non-ambiguous prompts, such as multiple
|
||||
input prompts, multimask_output=False can give better results.
|
||||
return_logits (bool): If true, returns un-thresholded masks logits
|
||||
instead of a binary mask.
|
||||
|
||||
Returns:
|
||||
(torch.Tensor): The output masks in BxCxHxW format, where C is the
|
||||
number of masks, and (H, W) is the original image size.
|
||||
(torch.Tensor): An array of shape BxC containing the model's
|
||||
predictions for the quality of each mask.
|
||||
(torch.Tensor): An array of shape BxCxHxW, where C is the number
|
||||
of masks and H=W=256. These low res logits can be passed to
|
||||
a subsequent iteration as mask input.
|
||||
"""
|
||||
if not self._is_image_set:
|
||||
raise RuntimeError(
|
||||
"An image must be set with .set_image(...) before mask prediction."
|
||||
)
|
||||
|
||||
if point_coords is not None:
|
||||
concat_points = (point_coords, point_labels)
|
||||
else:
|
||||
concat_points = None
|
||||
|
||||
# Embed prompts
|
||||
if boxes is not None:
|
||||
box_coords = boxes.reshape(-1, 2, 2)
|
||||
box_labels = torch.tensor([[2, 3]], dtype=torch.int, device=boxes.device)
|
||||
box_labels = box_labels.repeat(boxes.size(0), 1)
|
||||
# we merge "boxes" and "points" into a single "concat_points" input (where
|
||||
# boxes are added at the beginning) to sam_prompt_encoder
|
||||
if concat_points is not None:
|
||||
concat_coords = torch.cat([box_coords, concat_points[0]], dim=1)
|
||||
concat_labels = torch.cat([box_labels, concat_points[1]], dim=1)
|
||||
concat_points = (concat_coords, concat_labels)
|
||||
else:
|
||||
concat_points = (box_coords, box_labels)
|
||||
|
||||
sparse_embeddings, dense_embeddings = self.model.sam_prompt_encoder(
|
||||
points=concat_points,
|
||||
boxes=None,
|
||||
masks=mask_input,
|
||||
)
|
||||
|
||||
# Predict masks
|
||||
batched_mode = (
|
||||
concat_points is not None and concat_points[0].shape[0] > 1
|
||||
) # multi object prediction
|
||||
high_res_features = [
|
||||
feat_level[img_idx].unsqueeze(0)
|
||||
for feat_level in self._features["high_res_feats"]
|
||||
]
|
||||
low_res_masks, iou_predictions, _, _ = self.model.sam_mask_decoder(
|
||||
image_embeddings=self._features["image_embed"][img_idx].unsqueeze(0),
|
||||
image_pe=self.model.sam_prompt_encoder.get_dense_pe(),
|
||||
sparse_prompt_embeddings=sparse_embeddings,
|
||||
dense_prompt_embeddings=dense_embeddings,
|
||||
multimask_output=multimask_output,
|
||||
repeat_image=batched_mode,
|
||||
high_res_features=high_res_features,
|
||||
)
|
||||
|
||||
# Upscale the masks to the original image resolution
|
||||
masks = self._transforms.postprocess_masks(
|
||||
low_res_masks, self._orig_hw[img_idx]
|
||||
)
|
||||
low_res_masks = torch.clamp(low_res_masks, -32.0, 32.0)
|
||||
if not return_logits:
|
||||
masks = masks > self.mask_threshold
|
||||
|
||||
return masks, iou_predictions, low_res_masks
|
||||
|
||||
def get_image_embedding(self) -> torch.Tensor:
|
||||
"""
|
||||
Returns the image embeddings for the currently set image, with
|
||||
shape 1xCxHxW, where C is the embedding dimension and (H,W) are
|
||||
the embedding spatial dimension of SAM (typically C=256, H=W=64).
|
||||
"""
|
||||
if not self._is_image_set:
|
||||
raise RuntimeError(
|
||||
"An image must be set with .set_image(...) to generate an embedding."
|
||||
)
|
||||
assert (
|
||||
self._features is not None
|
||||
), "Features must exist if an image has been set."
|
||||
return self._features["image_embed"]
|
||||
|
||||
@property
|
||||
def device(self) -> torch.device:
|
||||
return self.model.device
|
||||
|
||||
def reset_predictor(self) -> None:
|
||||
"""
|
||||
Resets the image embeddings and other state variables.
|
||||
"""
|
||||
self._is_image_set = False
|
||||
self._features = None
|
||||
self._orig_hw = None
|
||||
self._is_batch = False
|
||||
@@ -0,0 +1,984 @@
|
||||
# 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.
|
||||
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
from model.segment_anything_2.sam2.modeling.sam2_base import NO_OBJ_SCORE, SAM2Base
|
||||
from model.segment_anything_2.sam2.utils.misc import concat_points, fill_holes_in_mask_scores, load_video_frames
|
||||
|
||||
|
||||
class SAM2VideoPredictor(SAM2Base):
|
||||
"""The predictor class to handle user interactions and manage inference states."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
fill_hole_area=0,
|
||||
# whether to apply non-overlapping constraints on the output object masks
|
||||
non_overlap_masks=False,
|
||||
# whether to clear non-conditioning memory of the surrounding frames (which may contain outdated information) after adding correction clicks;
|
||||
# note that this would only apply to *single-object tracking* unless `clear_non_cond_mem_for_multi_obj` is also set to True)
|
||||
clear_non_cond_mem_around_input=False,
|
||||
# whether to also clear non-conditioning memory of the surrounding frames (only effective when `clear_non_cond_mem_around_input` is True).
|
||||
clear_non_cond_mem_for_multi_obj=False,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
self.fill_hole_area = fill_hole_area
|
||||
self.non_overlap_masks = non_overlap_masks
|
||||
self.clear_non_cond_mem_around_input = clear_non_cond_mem_around_input
|
||||
self.clear_non_cond_mem_for_multi_obj = clear_non_cond_mem_for_multi_obj
|
||||
|
||||
@torch.inference_mode()
|
||||
def init_state(
|
||||
self,
|
||||
video_path,
|
||||
offload_video_to_cpu=False,
|
||||
offload_state_to_cpu=False,
|
||||
async_loading_frames=False,
|
||||
):
|
||||
"""Initialize a inference state."""
|
||||
images, video_height, video_width = load_video_frames(
|
||||
video_path=video_path,
|
||||
image_size=self.image_size,
|
||||
offload_video_to_cpu=offload_video_to_cpu,
|
||||
async_loading_frames=async_loading_frames,
|
||||
)
|
||||
inference_state = {}
|
||||
inference_state["images"] = images
|
||||
inference_state["num_frames"] = len(images)
|
||||
# whether to offload the video frames to CPU memory
|
||||
# turning on this option saves the GPU memory with only a very small overhead
|
||||
inference_state["offload_video_to_cpu"] = offload_video_to_cpu
|
||||
# whether to offload the inference state to CPU memory
|
||||
# turning on this option saves the GPU memory at the cost of a lower tracking fps
|
||||
# (e.g. in a test case of 768x768 model, fps dropped from 27 to 24 when tracking one object
|
||||
# and from 24 to 21 when tracking two objects)
|
||||
inference_state["offload_state_to_cpu"] = offload_state_to_cpu
|
||||
# the original video height and width, used for resizing final output scores
|
||||
inference_state["video_height"] = video_height
|
||||
inference_state["video_width"] = video_width
|
||||
inference_state["device"] = torch.device("cuda")
|
||||
if offload_state_to_cpu:
|
||||
inference_state["storage_device"] = torch.device("cpu")
|
||||
else:
|
||||
inference_state["storage_device"] = torch.device("cuda")
|
||||
# inputs on each frame
|
||||
inference_state["point_inputs_per_obj"] = {}
|
||||
inference_state["mask_inputs_per_obj"] = {}
|
||||
# visual features on a small number of recently visited frames for quick interactions
|
||||
inference_state["cached_features"] = {}
|
||||
# values that don't change across frames (so we only need to hold one copy of them)
|
||||
inference_state["constants"] = {}
|
||||
# mapping between client-side object id and model-side object index
|
||||
inference_state["obj_id_to_idx"] = OrderedDict()
|
||||
inference_state["obj_idx_to_id"] = OrderedDict()
|
||||
inference_state["obj_ids"] = []
|
||||
# A storage to hold the model's tracking results and states on each frame
|
||||
inference_state["output_dict"] = {
|
||||
"cond_frame_outputs": {}, # dict containing {frame_idx: <out>}
|
||||
"non_cond_frame_outputs": {}, # dict containing {frame_idx: <out>}
|
||||
}
|
||||
# Slice (view) of each object tracking results, sharing the same memory with "output_dict"
|
||||
inference_state["output_dict_per_obj"] = {}
|
||||
# A temporary storage to hold new outputs when user interact with a frame
|
||||
# to add clicks or mask (it's merged into "output_dict" before propagation starts)
|
||||
inference_state["temp_output_dict_per_obj"] = {}
|
||||
# Frames that already holds consolidated outputs from click or mask inputs
|
||||
# (we directly use their consolidated outputs during tracking)
|
||||
inference_state["consolidated_frame_inds"] = {
|
||||
"cond_frame_outputs": set(), # set containing frame indices
|
||||
"non_cond_frame_outputs": set(), # set containing frame indices
|
||||
}
|
||||
# metadata for each tracking frame (e.g. which direction it's tracked)
|
||||
inference_state["tracking_has_started"] = False
|
||||
inference_state["frames_already_tracked"] = {}
|
||||
# Warm up the visual backbone and cache the image feature on frame 0
|
||||
self._get_image_feature(inference_state, frame_idx=0, batch_size=1)
|
||||
return inference_state
|
||||
|
||||
def _obj_id_to_idx(self, inference_state, obj_id):
|
||||
"""Map client-side object id to model-side object index."""
|
||||
obj_idx = inference_state["obj_id_to_idx"].get(obj_id, None)
|
||||
if obj_idx is not None:
|
||||
return obj_idx
|
||||
|
||||
# This is a new object id not sent to the server before. We only allow adding
|
||||
# new objects *before* the tracking starts.
|
||||
allow_new_object = not inference_state["tracking_has_started"]
|
||||
if allow_new_object:
|
||||
# get the next object slot
|
||||
obj_idx = len(inference_state["obj_id_to_idx"])
|
||||
inference_state["obj_id_to_idx"][obj_id] = obj_idx
|
||||
inference_state["obj_idx_to_id"][obj_idx] = obj_id
|
||||
inference_state["obj_ids"] = list(inference_state["obj_id_to_idx"])
|
||||
# set up input and output structures for this object
|
||||
inference_state["point_inputs_per_obj"][obj_idx] = {}
|
||||
inference_state["mask_inputs_per_obj"][obj_idx] = {}
|
||||
inference_state["output_dict_per_obj"][obj_idx] = {
|
||||
"cond_frame_outputs": {}, # dict containing {frame_idx: <out>}
|
||||
"non_cond_frame_outputs": {}, # dict containing {frame_idx: <out>}
|
||||
}
|
||||
inference_state["temp_output_dict_per_obj"][obj_idx] = {
|
||||
"cond_frame_outputs": {}, # dict containing {frame_idx: <out>}
|
||||
"non_cond_frame_outputs": {}, # dict containing {frame_idx: <out>}
|
||||
}
|
||||
return obj_idx
|
||||
else:
|
||||
raise RuntimeError(
|
||||
f"Cannot add new object id {obj_id} after tracking starts. "
|
||||
f"All existing object ids: {inference_state['obj_ids']}. "
|
||||
f"Please call 'reset_state' to restart from scratch."
|
||||
)
|
||||
|
||||
def _obj_idx_to_id(self, inference_state, obj_idx):
|
||||
"""Map model-side object index to client-side object id."""
|
||||
return inference_state["obj_idx_to_id"][obj_idx]
|
||||
|
||||
def _get_obj_num(self, inference_state):
|
||||
"""Get the total number of unique object ids received so far in this session."""
|
||||
return len(inference_state["obj_idx_to_id"])
|
||||
|
||||
@torch.inference_mode()
|
||||
def add_new_points(
|
||||
self,
|
||||
inference_state,
|
||||
frame_idx,
|
||||
obj_id,
|
||||
points,
|
||||
labels,
|
||||
clear_old_points=True,
|
||||
normalize_coords=True,
|
||||
):
|
||||
"""Add new points to a frame."""
|
||||
obj_idx = self._obj_id_to_idx(inference_state, obj_id)
|
||||
point_inputs_per_frame = inference_state["point_inputs_per_obj"][obj_idx]
|
||||
mask_inputs_per_frame = inference_state["mask_inputs_per_obj"][obj_idx]
|
||||
|
||||
if not isinstance(points, torch.Tensor):
|
||||
points = torch.tensor(points, dtype=torch.float32)
|
||||
if not isinstance(labels, torch.Tensor):
|
||||
labels = torch.tensor(labels, dtype=torch.int32)
|
||||
if points.dim() == 2:
|
||||
points = points.unsqueeze(0) # add batch dimension
|
||||
if labels.dim() == 1:
|
||||
labels = labels.unsqueeze(0) # add batch dimension
|
||||
if normalize_coords:
|
||||
video_H = inference_state["video_height"]
|
||||
video_W = inference_state["video_width"]
|
||||
points = points / torch.tensor([video_W, video_H]).to(points.device)
|
||||
# scale the (normalized) coordinates by the model's internal image size
|
||||
points = points * self.image_size
|
||||
points = points.to(inference_state["device"])
|
||||
labels = labels.to(inference_state["device"])
|
||||
|
||||
if not clear_old_points:
|
||||
point_inputs = point_inputs_per_frame.get(frame_idx, None)
|
||||
else:
|
||||
point_inputs = None
|
||||
point_inputs = concat_points(point_inputs, points, labels)
|
||||
|
||||
point_inputs_per_frame[frame_idx] = point_inputs
|
||||
mask_inputs_per_frame.pop(frame_idx, None)
|
||||
# If this frame hasn't been tracked before, we treat it as an initial conditioning
|
||||
# frame, meaning that the inputs points are to generate segments on this frame without
|
||||
# using any memory from other frames, like in SAM. Otherwise (if it has been tracked),
|
||||
# the input points will be used to correct the already tracked masks.
|
||||
is_init_cond_frame = frame_idx not in inference_state["frames_already_tracked"]
|
||||
# whether to track in reverse time order
|
||||
if is_init_cond_frame:
|
||||
reverse = False
|
||||
else:
|
||||
reverse = inference_state["frames_already_tracked"][frame_idx]["reverse"]
|
||||
obj_output_dict = inference_state["output_dict_per_obj"][obj_idx]
|
||||
obj_temp_output_dict = inference_state["temp_output_dict_per_obj"][obj_idx]
|
||||
# Add a frame to conditioning output if it's an initial conditioning frame or
|
||||
# if the model sees all frames receiving clicks/mask as conditioning frames.
|
||||
is_cond = is_init_cond_frame or self.add_all_frames_to_correct_as_cond
|
||||
storage_key = "cond_frame_outputs" if is_cond else "non_cond_frame_outputs"
|
||||
|
||||
# Get any previously predicted mask logits on this object and feed it along with
|
||||
# the new clicks into the SAM mask decoder.
|
||||
prev_sam_mask_logits = None
|
||||
# lookup temporary output dict first, which contains the most recent output
|
||||
# (if not found, then lookup conditioning and non-conditioning frame output)
|
||||
prev_out = obj_temp_output_dict[storage_key].get(frame_idx)
|
||||
if prev_out is None:
|
||||
prev_out = obj_output_dict["cond_frame_outputs"].get(frame_idx)
|
||||
if prev_out is None:
|
||||
prev_out = obj_output_dict["non_cond_frame_outputs"].get(frame_idx)
|
||||
|
||||
if prev_out is not None and prev_out["pred_masks"] is not None:
|
||||
prev_sam_mask_logits = prev_out["pred_masks"].cuda(non_blocking=True)
|
||||
# Clamp the scale of prev_sam_mask_logits to avoid rare numerical issues.
|
||||
prev_sam_mask_logits = torch.clamp(prev_sam_mask_logits, -32.0, 32.0)
|
||||
current_out, _ = self._run_single_frame_inference(
|
||||
inference_state=inference_state,
|
||||
output_dict=obj_output_dict, # run on the slice of a single object
|
||||
frame_idx=frame_idx,
|
||||
batch_size=1, # run on the slice of a single object
|
||||
is_init_cond_frame=is_init_cond_frame,
|
||||
point_inputs=point_inputs,
|
||||
mask_inputs=None,
|
||||
reverse=reverse,
|
||||
# Skip the memory encoder when adding clicks or mask. We execute the memory encoder
|
||||
# at the beginning of `propagate_in_video` (after user finalize their clicks). This
|
||||
# allows us to enforce non-overlapping constraints on all objects before encoding
|
||||
# them into memory.
|
||||
run_mem_encoder=False,
|
||||
prev_sam_mask_logits=prev_sam_mask_logits,
|
||||
)
|
||||
# Add the output to the output dict (to be used as future memory)
|
||||
obj_temp_output_dict[storage_key][frame_idx] = current_out
|
||||
|
||||
# Resize the output mask to the original video resolution
|
||||
obj_ids = inference_state["obj_ids"]
|
||||
consolidated_out = self._consolidate_temp_output_across_obj(
|
||||
inference_state,
|
||||
frame_idx,
|
||||
is_cond=is_cond,
|
||||
run_mem_encoder=False,
|
||||
consolidate_at_video_res=True,
|
||||
)
|
||||
_, video_res_masks = self._get_orig_video_res_output(
|
||||
inference_state, consolidated_out["pred_masks_video_res"]
|
||||
)
|
||||
return frame_idx, obj_ids, video_res_masks
|
||||
|
||||
@torch.inference_mode()
|
||||
def add_new_mask(
|
||||
self,
|
||||
inference_state,
|
||||
frame_idx,
|
||||
obj_id,
|
||||
mask,
|
||||
):
|
||||
"""Add new mask to a frame."""
|
||||
obj_idx = self._obj_id_to_idx(inference_state, obj_id)
|
||||
point_inputs_per_frame = inference_state["point_inputs_per_obj"][obj_idx]
|
||||
mask_inputs_per_frame = inference_state["mask_inputs_per_obj"][obj_idx]
|
||||
|
||||
if not isinstance(mask, torch.Tensor):
|
||||
mask = torch.tensor(mask, dtype=torch.bool)
|
||||
assert mask.dim() == 2
|
||||
mask_H, mask_W = mask.shape
|
||||
mask_inputs_orig = mask[None, None] # add batch and channel dimension
|
||||
mask_inputs_orig = mask_inputs_orig.float().to(inference_state["device"])
|
||||
|
||||
# resize the mask if it doesn't match the model's image size
|
||||
if mask_H != self.image_size or mask_W != self.image_size:
|
||||
mask_inputs = torch.nn.functional.interpolate(
|
||||
mask_inputs_orig,
|
||||
size=(self.image_size, self.image_size),
|
||||
align_corners=False,
|
||||
mode="bilinear",
|
||||
antialias=True, # use antialias for downsampling
|
||||
)
|
||||
mask_inputs = (mask_inputs >= 0.5).float()
|
||||
else:
|
||||
mask_inputs = mask_inputs_orig
|
||||
|
||||
mask_inputs_per_frame[frame_idx] = mask_inputs
|
||||
point_inputs_per_frame.pop(frame_idx, None)
|
||||
# If this frame hasn't been tracked before, we treat it as an initial conditioning
|
||||
# frame, meaning that the inputs points are to generate segments on this frame without
|
||||
# using any memory from other frames, like in SAM. Otherwise (if it has been tracked),
|
||||
# the input points will be used to correct the already tracked masks.
|
||||
is_init_cond_frame = frame_idx not in inference_state["frames_already_tracked"]
|
||||
# whether to track in reverse time order
|
||||
if is_init_cond_frame:
|
||||
reverse = False
|
||||
else:
|
||||
reverse = inference_state["frames_already_tracked"][frame_idx]["reverse"]
|
||||
obj_output_dict = inference_state["output_dict_per_obj"][obj_idx]
|
||||
obj_temp_output_dict = inference_state["temp_output_dict_per_obj"][obj_idx]
|
||||
# Add a frame to conditioning output if it's an initial conditioning frame or
|
||||
# if the model sees all frames receiving clicks/mask as conditioning frames.
|
||||
is_cond = is_init_cond_frame or self.add_all_frames_to_correct_as_cond
|
||||
storage_key = "cond_frame_outputs" if is_cond else "non_cond_frame_outputs"
|
||||
|
||||
current_out, _ = self._run_single_frame_inference(
|
||||
inference_state=inference_state,
|
||||
output_dict=obj_output_dict, # run on the slice of a single object
|
||||
frame_idx=frame_idx,
|
||||
batch_size=1, # run on the slice of a single object
|
||||
is_init_cond_frame=is_init_cond_frame,
|
||||
point_inputs=None,
|
||||
mask_inputs=mask_inputs,
|
||||
reverse=reverse,
|
||||
# Skip the memory encoder when adding clicks or mask. We execute the memory encoder
|
||||
# at the beginning of `propagate_in_video` (after user finalize their clicks). This
|
||||
# allows us to enforce non-overlapping constraints on all objects before encoding
|
||||
# them into memory.
|
||||
run_mem_encoder=False,
|
||||
)
|
||||
# Add the output to the output dict (to be used as future memory)
|
||||
obj_temp_output_dict[storage_key][frame_idx] = current_out
|
||||
|
||||
# Resize the output mask to the original video resolution
|
||||
obj_ids = inference_state["obj_ids"]
|
||||
consolidated_out = self._consolidate_temp_output_across_obj(
|
||||
inference_state,
|
||||
frame_idx,
|
||||
is_cond=is_cond,
|
||||
run_mem_encoder=False,
|
||||
consolidate_at_video_res=True,
|
||||
)
|
||||
_, video_res_masks = self._get_orig_video_res_output(
|
||||
inference_state, consolidated_out["pred_masks_video_res"]
|
||||
)
|
||||
return frame_idx, obj_ids, video_res_masks
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def add_new_text(
|
||||
self,
|
||||
inference_state,
|
||||
frame_idx,
|
||||
obj_id,
|
||||
text,
|
||||
clear_old_points=True,
|
||||
normalize_coords=True,
|
||||
):
|
||||
"""Add new text to a frame."""
|
||||
obj_idx = self._obj_id_to_idx(inference_state, obj_id)
|
||||
point_inputs_per_frame = inference_state["point_inputs_per_obj"][obj_idx]
|
||||
mask_inputs_per_frame = inference_state["mask_inputs_per_obj"][obj_idx]
|
||||
|
||||
mask_inputs_per_frame.pop(frame_idx, None)
|
||||
# If this frame hasn't been tracked before, we treat it as an initial conditioning
|
||||
# frame, meaning that the inputs points are to generate segments on this frame without
|
||||
# using any memory from other frames, like in SAM. Otherwise (if it has been tracked),
|
||||
# the input points will be used to correct the already tracked masks.
|
||||
is_init_cond_frame = frame_idx not in inference_state["frames_already_tracked"]
|
||||
# whether to track in reverse time order
|
||||
if is_init_cond_frame:
|
||||
reverse = False
|
||||
else:
|
||||
reverse = inference_state["frames_already_tracked"][frame_idx]["reverse"]
|
||||
obj_output_dict = inference_state["output_dict_per_obj"][obj_idx]
|
||||
obj_temp_output_dict = inference_state["temp_output_dict_per_obj"][obj_idx]
|
||||
# Add a frame to conditioning output if it's an initial conditioning frame or
|
||||
# if the model sees all frames receiving clicks/mask as conditioning frames.
|
||||
is_cond = is_init_cond_frame or self.add_all_frames_to_correct_as_cond
|
||||
storage_key = "cond_frame_outputs" if is_cond else "non_cond_frame_outputs"
|
||||
|
||||
# Get any previously predicted mask logits on this object and feed it along with
|
||||
# the new clicks into the SAM mask decoder.
|
||||
prev_sam_mask_logits = None
|
||||
# lookup temporary output dict first, which contains the most recent output
|
||||
# (if not found, then lookup conditioning and non-conditioning frame output)
|
||||
prev_out = obj_temp_output_dict[storage_key].get(frame_idx)
|
||||
if prev_out is None:
|
||||
prev_out = obj_output_dict["cond_frame_outputs"].get(frame_idx)
|
||||
if prev_out is None:
|
||||
prev_out = obj_output_dict["non_cond_frame_outputs"].get(frame_idx)
|
||||
|
||||
if prev_out is not None and prev_out["pred_masks"] is not None:
|
||||
prev_sam_mask_logits = prev_out["pred_masks"].cuda(non_blocking=True)
|
||||
# Clamp the scale of prev_sam_mask_logits to avoid rare numerical issues.
|
||||
prev_sam_mask_logits = torch.clamp(prev_sam_mask_logits, -32.0, 32.0)
|
||||
current_out, _ = self._run_single_frame_inference(
|
||||
inference_state=inference_state,
|
||||
output_dict=obj_output_dict, # run on the slice of a single object
|
||||
frame_idx=frame_idx,
|
||||
batch_size=1, # run on the slice of a single object
|
||||
is_init_cond_frame=is_init_cond_frame,
|
||||
point_inputs=None,
|
||||
mask_inputs=None,
|
||||
reverse=reverse,
|
||||
# Skip the memory encoder when adding clicks or mask. We execute the memory encoder
|
||||
# at the beginning of `propagate_in_video` (after user finalize their clicks). This
|
||||
# allows us to enforce non-overlapping constraints on all objects before encoding
|
||||
# them into memory.
|
||||
run_mem_encoder=False,
|
||||
prev_sam_mask_logits=prev_sam_mask_logits,
|
||||
text_inputs=text
|
||||
)
|
||||
# Add the output to the output dict (to be used as future memory)
|
||||
obj_temp_output_dict[storage_key][frame_idx] = current_out
|
||||
|
||||
# Resize the output mask to the original video resolution
|
||||
obj_ids = inference_state["obj_ids"]
|
||||
consolidated_out = self._consolidate_temp_output_across_obj(
|
||||
inference_state,
|
||||
frame_idx,
|
||||
is_cond=is_cond,
|
||||
run_mem_encoder=False,
|
||||
consolidate_at_video_res=True,
|
||||
)
|
||||
_, video_res_masks = self._get_orig_video_res_output(
|
||||
inference_state, consolidated_out["pred_masks_video_res"]
|
||||
)
|
||||
return frame_idx, obj_ids, video_res_masks
|
||||
|
||||
|
||||
def _get_orig_video_res_output(self, inference_state, any_res_masks):
|
||||
"""
|
||||
Resize the object scores to the original video resolution (video_res_masks)
|
||||
and apply non-overlapping constraints for final output.
|
||||
"""
|
||||
device = inference_state["device"]
|
||||
video_H = inference_state["video_height"]
|
||||
video_W = inference_state["video_width"]
|
||||
any_res_masks = any_res_masks.to(device, non_blocking=True)
|
||||
if any_res_masks.shape[-2:] == (video_H, video_W):
|
||||
video_res_masks = any_res_masks
|
||||
else:
|
||||
video_res_masks = torch.nn.functional.interpolate(
|
||||
any_res_masks,
|
||||
size=(video_H, video_W),
|
||||
mode="bilinear",
|
||||
align_corners=False,
|
||||
)
|
||||
if self.non_overlap_masks:
|
||||
video_res_masks = self._apply_non_overlapping_constraints(video_res_masks)
|
||||
return any_res_masks, video_res_masks
|
||||
|
||||
def _consolidate_temp_output_across_obj(
|
||||
self,
|
||||
inference_state,
|
||||
frame_idx,
|
||||
is_cond,
|
||||
run_mem_encoder,
|
||||
consolidate_at_video_res=False,
|
||||
):
|
||||
"""
|
||||
Consolidate the per-object temporary outputs in `temp_output_dict_per_obj` on
|
||||
a frame into a single output for all objects, including
|
||||
1) fill any missing objects either from `output_dict_per_obj` (if they exist in
|
||||
`output_dict_per_obj` for this frame) or leave them as placeholder values
|
||||
(if they don't exist in `output_dict_per_obj` for this frame);
|
||||
2) if specified, rerun memory encoder after apply non-overlapping constraints
|
||||
on the object scores.
|
||||
"""
|
||||
batch_size = self._get_obj_num(inference_state)
|
||||
storage_key = "cond_frame_outputs" if is_cond else "non_cond_frame_outputs"
|
||||
# Optionally, we allow consolidating the temporary outputs at the original
|
||||
# video resolution (to provide a better editing experience for mask prompts).
|
||||
if consolidate_at_video_res:
|
||||
assert not run_mem_encoder, "memory encoder cannot run at video resolution"
|
||||
consolidated_H = inference_state["video_height"]
|
||||
consolidated_W = inference_state["video_width"]
|
||||
consolidated_mask_key = "pred_masks_video_res"
|
||||
else:
|
||||
consolidated_H = consolidated_W = self.image_size // 4
|
||||
consolidated_mask_key = "pred_masks"
|
||||
|
||||
# Initialize `consolidated_out`. Its "maskmem_features" and "maskmem_pos_enc"
|
||||
# will be added when rerunning the memory encoder after applying non-overlapping
|
||||
# constraints to object scores. Its "pred_masks" are prefilled with a large
|
||||
# negative value (NO_OBJ_SCORE) to represent missing objects.
|
||||
consolidated_out = {
|
||||
"maskmem_features": None,
|
||||
"maskmem_pos_enc": None,
|
||||
consolidated_mask_key: torch.full(
|
||||
size=(batch_size, 1, consolidated_H, consolidated_W),
|
||||
fill_value=NO_OBJ_SCORE,
|
||||
dtype=torch.float32,
|
||||
device=inference_state["storage_device"],
|
||||
),
|
||||
"obj_ptr": torch.full(
|
||||
size=(batch_size, self.hidden_dim),
|
||||
fill_value=NO_OBJ_SCORE,
|
||||
dtype=torch.float32,
|
||||
device=inference_state["device"],
|
||||
),
|
||||
}
|
||||
empty_mask_ptr = None
|
||||
for obj_idx in range(batch_size):
|
||||
obj_temp_output_dict = inference_state["temp_output_dict_per_obj"][obj_idx]
|
||||
obj_output_dict = inference_state["output_dict_per_obj"][obj_idx]
|
||||
out = obj_temp_output_dict[storage_key].get(frame_idx, None)
|
||||
# If the object doesn't appear in "temp_output_dict_per_obj" on this frame,
|
||||
# we fall back and look up its previous output in "output_dict_per_obj".
|
||||
# We look up both "cond_frame_outputs" and "non_cond_frame_outputs" in
|
||||
# "output_dict_per_obj" to find a previous output for this object.
|
||||
if out is None:
|
||||
out = obj_output_dict["cond_frame_outputs"].get(frame_idx, None)
|
||||
if out is None:
|
||||
out = obj_output_dict["non_cond_frame_outputs"].get(frame_idx, None)
|
||||
# If the object doesn't appear in "output_dict_per_obj" either, we skip it
|
||||
# and leave its mask scores to the default scores (i.e. the NO_OBJ_SCORE
|
||||
# placeholder above) and set its object pointer to be a dummy pointer.
|
||||
if out is None:
|
||||
# Fill in dummy object pointers for those objects without any inputs or
|
||||
# tracking outcomes on this frame (only do it under `run_mem_encoder=True`,
|
||||
# i.e. when we need to build the memory for tracking).
|
||||
if run_mem_encoder:
|
||||
if empty_mask_ptr is None:
|
||||
empty_mask_ptr = self._get_empty_mask_ptr(
|
||||
inference_state, frame_idx
|
||||
)
|
||||
# fill object pointer with a dummy pointer (based on an empty mask)
|
||||
consolidated_out["obj_ptr"][obj_idx : obj_idx + 1] = empty_mask_ptr
|
||||
continue
|
||||
# Add the temporary object output mask to consolidated output mask
|
||||
obj_mask = out["pred_masks"]
|
||||
consolidated_pred_masks = consolidated_out[consolidated_mask_key]
|
||||
if obj_mask.shape[-2:] == consolidated_pred_masks.shape[-2:]:
|
||||
consolidated_pred_masks[obj_idx : obj_idx + 1] = obj_mask
|
||||
else:
|
||||
# Resize first if temporary object mask has a different resolution
|
||||
resized_obj_mask = torch.nn.functional.interpolate(
|
||||
obj_mask,
|
||||
size=consolidated_pred_masks.shape[-2:],
|
||||
mode="bilinear",
|
||||
align_corners=False,
|
||||
)
|
||||
consolidated_pred_masks[obj_idx : obj_idx + 1] = resized_obj_mask
|
||||
consolidated_out["obj_ptr"][obj_idx : obj_idx + 1] = out["obj_ptr"]
|
||||
|
||||
# Optionally, apply non-overlapping constraints on the consolidated scores
|
||||
# and rerun the memory encoder
|
||||
if run_mem_encoder:
|
||||
device = inference_state["device"]
|
||||
high_res_masks = torch.nn.functional.interpolate(
|
||||
consolidated_out["pred_masks"].to(device, non_blocking=True),
|
||||
size=(self.image_size, self.image_size),
|
||||
mode="bilinear",
|
||||
align_corners=False,
|
||||
)
|
||||
if self.non_overlap_masks_for_mem_enc:
|
||||
high_res_masks = self._apply_non_overlapping_constraints(high_res_masks)
|
||||
maskmem_features, maskmem_pos_enc = self._run_memory_encoder(
|
||||
inference_state=inference_state,
|
||||
frame_idx=frame_idx,
|
||||
batch_size=batch_size,
|
||||
high_res_masks=high_res_masks,
|
||||
is_mask_from_pts=True, # these frames are what the user interacted with
|
||||
)
|
||||
consolidated_out["maskmem_features"] = maskmem_features
|
||||
consolidated_out["maskmem_pos_enc"] = maskmem_pos_enc
|
||||
|
||||
return consolidated_out
|
||||
|
||||
def _get_empty_mask_ptr(self, inference_state, frame_idx):
|
||||
"""Get a dummy object pointer based on an empty mask on the current frame."""
|
||||
# A dummy (empty) mask with a single object
|
||||
batch_size = 1
|
||||
mask_inputs = torch.zeros(
|
||||
(batch_size, 1, self.image_size, self.image_size),
|
||||
dtype=torch.float32,
|
||||
device=inference_state["device"],
|
||||
)
|
||||
|
||||
# Retrieve correct image features
|
||||
(
|
||||
_,
|
||||
_,
|
||||
current_vision_feats,
|
||||
current_vision_pos_embeds,
|
||||
feat_sizes,
|
||||
) = self._get_image_feature(inference_state, frame_idx, batch_size)
|
||||
|
||||
# Feed the empty mask and image feature above to get a dummy object pointer
|
||||
current_out = self.track_step(
|
||||
frame_idx=frame_idx,
|
||||
is_init_cond_frame=True,
|
||||
current_vision_feats=current_vision_feats,
|
||||
current_vision_pos_embeds=current_vision_pos_embeds,
|
||||
feat_sizes=feat_sizes,
|
||||
point_inputs=None,
|
||||
mask_inputs=mask_inputs,
|
||||
output_dict={},
|
||||
num_frames=inference_state["num_frames"],
|
||||
track_in_reverse=False,
|
||||
run_mem_encoder=False,
|
||||
prev_sam_mask_logits=None,
|
||||
)
|
||||
return current_out["obj_ptr"]
|
||||
|
||||
@torch.inference_mode()
|
||||
def propagate_in_video_preflight(self, inference_state):
|
||||
"""Prepare inference_state and consolidate temporary outputs before tracking."""
|
||||
# Tracking has started and we don't allow adding new objects until session is reset.
|
||||
inference_state["tracking_has_started"] = True
|
||||
batch_size = self._get_obj_num(inference_state)
|
||||
|
||||
# Consolidate per-object temporary outputs in "temp_output_dict_per_obj" and
|
||||
# add them into "output_dict".
|
||||
temp_output_dict_per_obj = inference_state["temp_output_dict_per_obj"]
|
||||
output_dict = inference_state["output_dict"]
|
||||
# "consolidated_frame_inds" contains indices of those frames where consolidated
|
||||
# temporary outputs have been added (either in this call or any previous calls
|
||||
# to `propagate_in_video_preflight`).
|
||||
consolidated_frame_inds = inference_state["consolidated_frame_inds"]
|
||||
for is_cond in [False, True]:
|
||||
# Separately consolidate conditioning and non-conditioning temp outptus
|
||||
storage_key = "cond_frame_outputs" if is_cond else "non_cond_frame_outputs"
|
||||
# Find all the frames that contain temporary outputs for any objects
|
||||
# (these should be the frames that have just received clicks for mask inputs
|
||||
# via `add_new_points` or `add_new_mask`)
|
||||
temp_frame_inds = set()
|
||||
for obj_temp_output_dict in temp_output_dict_per_obj.values():
|
||||
temp_frame_inds.update(obj_temp_output_dict[storage_key].keys())
|
||||
consolidated_frame_inds[storage_key].update(temp_frame_inds)
|
||||
# consolidate the temprary output across all objects on this frame
|
||||
for frame_idx in temp_frame_inds:
|
||||
consolidated_out = self._consolidate_temp_output_across_obj(
|
||||
inference_state, frame_idx, is_cond=is_cond, run_mem_encoder=True
|
||||
)
|
||||
# merge them into "output_dict" and also create per-object slices
|
||||
output_dict[storage_key][frame_idx] = consolidated_out
|
||||
self._add_output_per_object(
|
||||
inference_state, frame_idx, consolidated_out, storage_key
|
||||
)
|
||||
clear_non_cond_mem = self.clear_non_cond_mem_around_input and (
|
||||
self.clear_non_cond_mem_for_multi_obj or batch_size <= 1
|
||||
)
|
||||
if clear_non_cond_mem:
|
||||
# clear non-conditioning memory of the surrounding frames
|
||||
self._clear_non_cond_mem_around_input(inference_state, frame_idx)
|
||||
|
||||
# clear temporary outputs in `temp_output_dict_per_obj`
|
||||
for obj_temp_output_dict in temp_output_dict_per_obj.values():
|
||||
obj_temp_output_dict[storage_key].clear()
|
||||
|
||||
# edge case: if an output is added to "cond_frame_outputs", we remove any prior
|
||||
# output on the same frame in "non_cond_frame_outputs"
|
||||
for frame_idx in output_dict["cond_frame_outputs"]:
|
||||
output_dict["non_cond_frame_outputs"].pop(frame_idx, None)
|
||||
for obj_output_dict in inference_state["output_dict_per_obj"].values():
|
||||
for frame_idx in obj_output_dict["cond_frame_outputs"]:
|
||||
obj_output_dict["non_cond_frame_outputs"].pop(frame_idx, None)
|
||||
for frame_idx in consolidated_frame_inds["cond_frame_outputs"]:
|
||||
assert frame_idx in output_dict["cond_frame_outputs"]
|
||||
consolidated_frame_inds["non_cond_frame_outputs"].discard(frame_idx)
|
||||
|
||||
# Make sure that the frame indices in "consolidated_frame_inds" are exactly those frames
|
||||
# with either points or mask inputs (which should be true under a correct workflow).
|
||||
# all_consolidated_frame_inds = (
|
||||
# consolidated_frame_inds["cond_frame_outputs"]
|
||||
# | consolidated_frame_inds["non_cond_frame_outputs"]
|
||||
# )
|
||||
# input_frames_inds = set()
|
||||
# for point_inputs_per_frame in inference_state["point_inputs_per_obj"].values():
|
||||
# input_frames_inds.update(point_inputs_per_frame.keys())
|
||||
# for mask_inputs_per_frame in inference_state["mask_inputs_per_obj"].values():
|
||||
# input_frames_inds.update(mask_inputs_per_frame.keys())
|
||||
# assert all_consolidated_frame_inds == input_frames_inds
|
||||
|
||||
@torch.inference_mode()
|
||||
def propagate_in_video(
|
||||
self,
|
||||
inference_state,
|
||||
start_frame_idx=None,
|
||||
max_frame_num_to_track=None,
|
||||
reverse=False,
|
||||
):
|
||||
"""Propagate the input points across frames to track in the entire video."""
|
||||
self.propagate_in_video_preflight(inference_state)
|
||||
|
||||
output_dict = inference_state["output_dict"]
|
||||
consolidated_frame_inds = inference_state["consolidated_frame_inds"]
|
||||
obj_ids = inference_state["obj_ids"]
|
||||
num_frames = inference_state["num_frames"]
|
||||
batch_size = self._get_obj_num(inference_state)
|
||||
if len(output_dict["cond_frame_outputs"]) == 0:
|
||||
raise RuntimeError("No points are provided; please add points first")
|
||||
clear_non_cond_mem = self.clear_non_cond_mem_around_input and (
|
||||
self.clear_non_cond_mem_for_multi_obj or batch_size <= 1
|
||||
)
|
||||
|
||||
# set start index, end index, and processing order
|
||||
if start_frame_idx is None:
|
||||
# default: start from the earliest frame with input points
|
||||
start_frame_idx = min(output_dict["cond_frame_outputs"])
|
||||
if max_frame_num_to_track is None:
|
||||
# default: track all the frames in the video
|
||||
max_frame_num_to_track = num_frames
|
||||
if reverse:
|
||||
end_frame_idx = max(start_frame_idx - max_frame_num_to_track, 0)
|
||||
if start_frame_idx > 0:
|
||||
processing_order = range(start_frame_idx, end_frame_idx - 1, -1)
|
||||
else:
|
||||
processing_order = [] # skip reverse tracking if starting from frame 0
|
||||
else:
|
||||
end_frame_idx = min(
|
||||
start_frame_idx + max_frame_num_to_track, num_frames - 1
|
||||
)
|
||||
processing_order = range(start_frame_idx, end_frame_idx + 1)
|
||||
|
||||
for frame_idx in tqdm(processing_order, desc="propagate in video"):
|
||||
# We skip those frames already in consolidated outputs (these are frames
|
||||
# that received input clicks or mask). Note that we cannot directly run
|
||||
# batched forward on them via `_run_single_frame_inference` because the
|
||||
# number of clicks on each object might be different.
|
||||
if frame_idx in consolidated_frame_inds["cond_frame_outputs"]:
|
||||
storage_key = "cond_frame_outputs"
|
||||
current_out = output_dict[storage_key][frame_idx]
|
||||
pred_masks = current_out["pred_masks"]
|
||||
if clear_non_cond_mem:
|
||||
# clear non-conditioning memory of the surrounding frames
|
||||
self._clear_non_cond_mem_around_input(inference_state, frame_idx)
|
||||
elif frame_idx in consolidated_frame_inds["non_cond_frame_outputs"]:
|
||||
storage_key = "non_cond_frame_outputs"
|
||||
current_out = output_dict[storage_key][frame_idx]
|
||||
pred_masks = current_out["pred_masks"]
|
||||
else:
|
||||
storage_key = "non_cond_frame_outputs"
|
||||
current_out, pred_masks = self._run_single_frame_inference(
|
||||
inference_state=inference_state,
|
||||
output_dict=output_dict,
|
||||
frame_idx=frame_idx,
|
||||
batch_size=batch_size,
|
||||
is_init_cond_frame=False,
|
||||
point_inputs=None,
|
||||
mask_inputs=None,
|
||||
reverse=reverse,
|
||||
run_mem_encoder=True,
|
||||
)
|
||||
output_dict[storage_key][frame_idx] = current_out
|
||||
# Create slices of per-object outputs for subsequent interaction with each
|
||||
# individual object after tracking.
|
||||
self._add_output_per_object(
|
||||
inference_state, frame_idx, current_out, storage_key
|
||||
)
|
||||
inference_state["frames_already_tracked"][frame_idx] = {"reverse": reverse}
|
||||
|
||||
# Resize the output mask to the original video resolution (we directly use
|
||||
# the mask scores on GPU for output to avoid any CPU conversion in between)
|
||||
_, video_res_masks = self._get_orig_video_res_output(
|
||||
inference_state, pred_masks
|
||||
)
|
||||
yield frame_idx, obj_ids, video_res_masks
|
||||
|
||||
def _add_output_per_object(
|
||||
self, inference_state, frame_idx, current_out, storage_key
|
||||
):
|
||||
"""
|
||||
Split a multi-object output into per-object output slices and add them into
|
||||
`output_dict_per_obj`. The resulting slices share the same tensor storage.
|
||||
"""
|
||||
maskmem_features = current_out["maskmem_features"]
|
||||
assert maskmem_features is None or isinstance(maskmem_features, torch.Tensor)
|
||||
|
||||
maskmem_pos_enc = current_out["maskmem_pos_enc"]
|
||||
assert maskmem_pos_enc is None or isinstance(maskmem_pos_enc, list)
|
||||
|
||||
output_dict_per_obj = inference_state["output_dict_per_obj"]
|
||||
for obj_idx, obj_output_dict in output_dict_per_obj.items():
|
||||
obj_slice = slice(obj_idx, obj_idx + 1)
|
||||
obj_out = {
|
||||
"maskmem_features": None,
|
||||
"maskmem_pos_enc": None,
|
||||
"pred_masks": current_out["pred_masks"][obj_slice],
|
||||
"obj_ptr": current_out["obj_ptr"][obj_slice],
|
||||
}
|
||||
if maskmem_features is not None:
|
||||
obj_out["maskmem_features"] = maskmem_features[obj_slice]
|
||||
if maskmem_pos_enc is not None:
|
||||
obj_out["maskmem_pos_enc"] = [x[obj_slice] for x in maskmem_pos_enc]
|
||||
obj_output_dict[storage_key][frame_idx] = obj_out
|
||||
|
||||
@torch.inference_mode()
|
||||
def reset_state(self, inference_state):
|
||||
"""Remove all input points or mask in all frames throughout the video."""
|
||||
self._reset_tracking_results(inference_state)
|
||||
# Remove all object ids
|
||||
inference_state["obj_id_to_idx"].clear()
|
||||
inference_state["obj_idx_to_id"].clear()
|
||||
inference_state["obj_ids"].clear()
|
||||
inference_state["point_inputs_per_obj"].clear()
|
||||
inference_state["mask_inputs_per_obj"].clear()
|
||||
inference_state["output_dict_per_obj"].clear()
|
||||
inference_state["temp_output_dict_per_obj"].clear()
|
||||
|
||||
def _reset_tracking_results(self, inference_state):
|
||||
"""Reset all tracking inputs and results across the videos."""
|
||||
for v in inference_state["point_inputs_per_obj"].values():
|
||||
v.clear()
|
||||
for v in inference_state["mask_inputs_per_obj"].values():
|
||||
v.clear()
|
||||
for v in inference_state["output_dict_per_obj"].values():
|
||||
v["cond_frame_outputs"].clear()
|
||||
v["non_cond_frame_outputs"].clear()
|
||||
for v in inference_state["temp_output_dict_per_obj"].values():
|
||||
v["cond_frame_outputs"].clear()
|
||||
v["non_cond_frame_outputs"].clear()
|
||||
inference_state["output_dict"]["cond_frame_outputs"].clear()
|
||||
inference_state["output_dict"]["non_cond_frame_outputs"].clear()
|
||||
inference_state["consolidated_frame_inds"]["cond_frame_outputs"].clear()
|
||||
inference_state["consolidated_frame_inds"]["non_cond_frame_outputs"].clear()
|
||||
inference_state["tracking_has_started"] = False
|
||||
inference_state["frames_already_tracked"].clear()
|
||||
|
||||
def _get_image_feature(self, inference_state, frame_idx, batch_size):
|
||||
"""Compute the image features on a given frame."""
|
||||
# Look up in the cache first
|
||||
image, backbone_out = inference_state["cached_features"].get(
|
||||
frame_idx, (None, None)
|
||||
)
|
||||
if backbone_out is None:
|
||||
# Cache miss -- we will run inference on a single image
|
||||
image = inference_state["images"][frame_idx].cuda().float().unsqueeze(0)
|
||||
backbone_out = self.forward_image(image)
|
||||
# Cache the most recent frame's feature (for repeated interactions with
|
||||
# a frame; we can use an LRU cache for more frames in the future).
|
||||
inference_state["cached_features"] = {frame_idx: (image, backbone_out)}
|
||||
|
||||
# expand the features to have the same dimension as the number of objects
|
||||
expanded_image = image.expand(batch_size, -1, -1, -1)
|
||||
expanded_backbone_out = {
|
||||
"backbone_fpn": backbone_out["backbone_fpn"].copy(),
|
||||
"vision_pos_enc": backbone_out["vision_pos_enc"].copy(),
|
||||
}
|
||||
for i, feat in enumerate(expanded_backbone_out["backbone_fpn"]):
|
||||
expanded_backbone_out["backbone_fpn"][i] = feat.expand(
|
||||
batch_size, -1, -1, -1
|
||||
)
|
||||
for i, pos in enumerate(expanded_backbone_out["vision_pos_enc"]):
|
||||
pos = pos.expand(batch_size, -1, -1, -1)
|
||||
expanded_backbone_out["vision_pos_enc"][i] = pos
|
||||
|
||||
features = self._prepare_backbone_features(expanded_backbone_out)
|
||||
features = (expanded_image,) + features
|
||||
return features
|
||||
|
||||
def _run_single_frame_inference(
|
||||
self,
|
||||
inference_state,
|
||||
output_dict,
|
||||
frame_idx,
|
||||
batch_size,
|
||||
is_init_cond_frame,
|
||||
point_inputs,
|
||||
mask_inputs,
|
||||
reverse,
|
||||
run_mem_encoder,
|
||||
prev_sam_mask_logits=None,
|
||||
text_inputs=None
|
||||
):
|
||||
"""Run tracking on a single frame based on current inputs and previous memory."""
|
||||
# Retrieve correct image features
|
||||
(
|
||||
_,
|
||||
_,
|
||||
current_vision_feats,
|
||||
current_vision_pos_embeds,
|
||||
feat_sizes,
|
||||
) = self._get_image_feature(inference_state, frame_idx, batch_size)
|
||||
|
||||
# point and mask should not appear as input simultaneously on the same frame
|
||||
assert point_inputs is None or mask_inputs is None
|
||||
current_out = self.track_step(
|
||||
frame_idx=frame_idx,
|
||||
is_init_cond_frame=is_init_cond_frame,
|
||||
current_vision_feats=current_vision_feats,
|
||||
current_vision_pos_embeds=current_vision_pos_embeds,
|
||||
feat_sizes=feat_sizes,
|
||||
point_inputs=point_inputs,
|
||||
mask_inputs=mask_inputs,
|
||||
output_dict=output_dict,
|
||||
num_frames=inference_state["num_frames"],
|
||||
track_in_reverse=reverse,
|
||||
run_mem_encoder=run_mem_encoder,
|
||||
prev_sam_mask_logits=prev_sam_mask_logits,
|
||||
text_inputs=text_inputs
|
||||
)
|
||||
|
||||
# optionally offload the output to CPU memory to save GPU space
|
||||
storage_device = inference_state["storage_device"]
|
||||
maskmem_features = current_out["maskmem_features"]
|
||||
if maskmem_features is not None:
|
||||
maskmem_features = maskmem_features.to(torch.bfloat16)
|
||||
maskmem_features = maskmem_features.to(storage_device, non_blocking=True)
|
||||
pred_masks_gpu = current_out["pred_masks"]
|
||||
# potentially fill holes in the predicted masks
|
||||
if self.fill_hole_area > 0:
|
||||
pred_masks_gpu = fill_holes_in_mask_scores(
|
||||
pred_masks_gpu, self.fill_hole_area
|
||||
)
|
||||
pred_masks = pred_masks_gpu.to(storage_device, non_blocking=True)
|
||||
# "maskmem_pos_enc" is the same across frames, so we only need to store one copy of it
|
||||
maskmem_pos_enc = self._get_maskmem_pos_enc(inference_state, current_out)
|
||||
# object pointer is a small tensor, so we always keep it on GPU memory for fast access
|
||||
obj_ptr = current_out["obj_ptr"]
|
||||
# make a compact version of this frame's output to reduce the state size
|
||||
compact_current_out = {
|
||||
"maskmem_features": maskmem_features,
|
||||
"maskmem_pos_enc": maskmem_pos_enc,
|
||||
"pred_masks": pred_masks,
|
||||
"obj_ptr": obj_ptr,
|
||||
}
|
||||
return compact_current_out, pred_masks_gpu
|
||||
|
||||
def _run_memory_encoder(
|
||||
self, inference_state, frame_idx, batch_size, high_res_masks, is_mask_from_pts
|
||||
):
|
||||
"""
|
||||
Run the memory encoder on `high_res_masks`. This is usually after applying
|
||||
non-overlapping constraints to object scores. Since their scores changed, their
|
||||
memory also need to be computed again with the memory encoder.
|
||||
"""
|
||||
# Retrieve correct image features
|
||||
_, _, current_vision_feats, _, feat_sizes = self._get_image_feature(
|
||||
inference_state, frame_idx, batch_size
|
||||
)
|
||||
maskmem_features, maskmem_pos_enc = self._encode_new_memory(
|
||||
current_vision_feats=current_vision_feats,
|
||||
feat_sizes=feat_sizes,
|
||||
pred_masks_high_res=high_res_masks,
|
||||
is_mask_from_pts=is_mask_from_pts,
|
||||
)
|
||||
|
||||
# optionally offload the output to CPU memory to save GPU space
|
||||
storage_device = inference_state["storage_device"]
|
||||
maskmem_features = maskmem_features.to(torch.bfloat16)
|
||||
maskmem_features = maskmem_features.to(storage_device, non_blocking=True)
|
||||
# "maskmem_pos_enc" is the same across frames, so we only need to store one copy of it
|
||||
maskmem_pos_enc = self._get_maskmem_pos_enc(
|
||||
inference_state, {"maskmem_pos_enc": maskmem_pos_enc}
|
||||
)
|
||||
return maskmem_features, maskmem_pos_enc
|
||||
|
||||
def _get_maskmem_pos_enc(self, inference_state, current_out):
|
||||
"""
|
||||
`maskmem_pos_enc` is the same across frames and objects, so we cache it as
|
||||
a constant in the inference session to reduce session storage size.
|
||||
"""
|
||||
model_constants = inference_state["constants"]
|
||||
# "out_maskmem_pos_enc" should be either a list of tensors or None
|
||||
out_maskmem_pos_enc = current_out["maskmem_pos_enc"]
|
||||
if out_maskmem_pos_enc is not None:
|
||||
if "maskmem_pos_enc" not in model_constants:
|
||||
assert isinstance(out_maskmem_pos_enc, list)
|
||||
# only take the slice for one object, since it's same across objects
|
||||
maskmem_pos_enc = [x[0:1].clone() for x in out_maskmem_pos_enc]
|
||||
model_constants["maskmem_pos_enc"] = maskmem_pos_enc
|
||||
else:
|
||||
maskmem_pos_enc = model_constants["maskmem_pos_enc"]
|
||||
# expand the cached maskmem_pos_enc to the actual batch size
|
||||
batch_size = out_maskmem_pos_enc[0].size(0)
|
||||
expanded_maskmem_pos_enc = [
|
||||
x.expand(batch_size, -1, -1, -1) for x in maskmem_pos_enc
|
||||
]
|
||||
else:
|
||||
expanded_maskmem_pos_enc = None
|
||||
return expanded_maskmem_pos_enc
|
||||
|
||||
def _clear_non_cond_mem_around_input(self, inference_state, frame_idx):
|
||||
"""
|
||||
Remove the non-conditioning memory around the input frame. When users provide
|
||||
correction clicks, the surrounding frames' non-conditioning memories can still
|
||||
contain outdated object appearance information and could confuse the model.
|
||||
|
||||
This method clears those non-conditioning memories surrounding the interacted
|
||||
frame to avoid giving the model both old and new information about the object.
|
||||
"""
|
||||
r = self.memory_temporal_stride_for_eval
|
||||
frame_idx_begin = frame_idx - r * self.num_maskmem
|
||||
frame_idx_end = frame_idx + r * self.num_maskmem
|
||||
output_dict = inference_state["output_dict"]
|
||||
non_cond_frame_outputs = output_dict["non_cond_frame_outputs"]
|
||||
for t in range(frame_idx_begin, frame_idx_end + 1):
|
||||
non_cond_frame_outputs.pop(t, None)
|
||||
for obj_output_dict in inference_state["output_dict_per_obj"].values():
|
||||
obj_output_dict["non_cond_frame_outputs"].pop(t, None)
|
||||
@@ -0,0 +1,5 @@
|
||||
# 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.
|
||||
@@ -0,0 +1,348 @@
|
||||
# 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 math
|
||||
from copy import deepcopy
|
||||
from itertools import product
|
||||
from typing import Any, Dict, Generator, ItemsView, List, Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
# Very lightly adapted from https://github.com/facebookresearch/segment-anything/blob/main/segment_anything/utils/amg.py
|
||||
|
||||
|
||||
class MaskData:
|
||||
"""
|
||||
A structure for storing masks and their related data in batched format.
|
||||
Implements basic filtering and concatenation.
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs) -> None:
|
||||
for v in kwargs.values():
|
||||
assert isinstance(
|
||||
v, (list, np.ndarray, torch.Tensor)
|
||||
), "MaskData only supports list, numpy arrays, and torch tensors."
|
||||
self._stats = dict(**kwargs)
|
||||
|
||||
def __setitem__(self, key: str, item: Any) -> None:
|
||||
assert isinstance(
|
||||
item, (list, np.ndarray, torch.Tensor)
|
||||
), "MaskData only supports list, numpy arrays, and torch tensors."
|
||||
self._stats[key] = item
|
||||
|
||||
def __delitem__(self, key: str) -> None:
|
||||
del self._stats[key]
|
||||
|
||||
def __getitem__(self, key: str) -> Any:
|
||||
return self._stats[key]
|
||||
|
||||
def items(self) -> ItemsView[str, Any]:
|
||||
return self._stats.items()
|
||||
|
||||
def filter(self, keep: torch.Tensor) -> None:
|
||||
for k, v in self._stats.items():
|
||||
if v is None:
|
||||
self._stats[k] = None
|
||||
elif isinstance(v, torch.Tensor):
|
||||
self._stats[k] = v[torch.as_tensor(keep, device=v.device)]
|
||||
elif isinstance(v, np.ndarray):
|
||||
self._stats[k] = v[keep.detach().cpu().numpy()]
|
||||
elif isinstance(v, list) and keep.dtype == torch.bool:
|
||||
self._stats[k] = [a for i, a in enumerate(v) if keep[i]]
|
||||
elif isinstance(v, list):
|
||||
self._stats[k] = [v[i] for i in keep]
|
||||
else:
|
||||
raise TypeError(f"MaskData key {k} has an unsupported type {type(v)}.")
|
||||
|
||||
def cat(self, new_stats: "MaskData") -> None:
|
||||
for k, v in new_stats.items():
|
||||
if k not in self._stats or self._stats[k] is None:
|
||||
self._stats[k] = deepcopy(v)
|
||||
elif isinstance(v, torch.Tensor):
|
||||
self._stats[k] = torch.cat([self._stats[k], v], dim=0)
|
||||
elif isinstance(v, np.ndarray):
|
||||
self._stats[k] = np.concatenate([self._stats[k], v], axis=0)
|
||||
elif isinstance(v, list):
|
||||
self._stats[k] = self._stats[k] + deepcopy(v)
|
||||
else:
|
||||
raise TypeError(f"MaskData key {k} has an unsupported type {type(v)}.")
|
||||
|
||||
def to_numpy(self) -> None:
|
||||
for k, v in self._stats.items():
|
||||
if isinstance(v, torch.Tensor):
|
||||
self._stats[k] = v.float().detach().cpu().numpy()
|
||||
|
||||
|
||||
def is_box_near_crop_edge(
|
||||
boxes: torch.Tensor, crop_box: List[int], orig_box: List[int], atol: float = 20.0
|
||||
) -> torch.Tensor:
|
||||
"""Filter masks at the edge of a crop, but not at the edge of the original image."""
|
||||
crop_box_torch = torch.as_tensor(crop_box, dtype=torch.float, device=boxes.device)
|
||||
orig_box_torch = torch.as_tensor(orig_box, dtype=torch.float, device=boxes.device)
|
||||
boxes = uncrop_boxes_xyxy(boxes, crop_box).float()
|
||||
near_crop_edge = torch.isclose(boxes, crop_box_torch[None, :], atol=atol, rtol=0)
|
||||
near_image_edge = torch.isclose(boxes, orig_box_torch[None, :], atol=atol, rtol=0)
|
||||
near_crop_edge = torch.logical_and(near_crop_edge, ~near_image_edge)
|
||||
return torch.any(near_crop_edge, dim=1)
|
||||
|
||||
|
||||
def box_xyxy_to_xywh(box_xyxy: torch.Tensor) -> torch.Tensor:
|
||||
box_xywh = deepcopy(box_xyxy)
|
||||
box_xywh[2] = box_xywh[2] - box_xywh[0]
|
||||
box_xywh[3] = box_xywh[3] - box_xywh[1]
|
||||
return box_xywh
|
||||
|
||||
|
||||
def batch_iterator(batch_size: int, *args) -> Generator[List[Any], None, None]:
|
||||
assert len(args) > 0 and all(
|
||||
len(a) == len(args[0]) for a in args
|
||||
), "Batched iteration must have inputs of all the same size."
|
||||
n_batches = len(args[0]) // batch_size + int(len(args[0]) % batch_size != 0)
|
||||
for b in range(n_batches):
|
||||
yield [arg[b * batch_size : (b + 1) * batch_size] for arg in args]
|
||||
|
||||
|
||||
def mask_to_rle_pytorch(tensor: torch.Tensor) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Encodes masks to an uncompressed RLE, in the format expected by
|
||||
pycoco tools.
|
||||
"""
|
||||
# Put in fortran order and flatten h,w
|
||||
b, h, w = tensor.shape
|
||||
tensor = tensor.permute(0, 2, 1).flatten(1)
|
||||
|
||||
# Compute change indices
|
||||
diff = tensor[:, 1:] ^ tensor[:, :-1]
|
||||
change_indices = diff.nonzero()
|
||||
|
||||
# Encode run length
|
||||
out = []
|
||||
for i in range(b):
|
||||
cur_idxs = change_indices[change_indices[:, 0] == i, 1]
|
||||
cur_idxs = torch.cat(
|
||||
[
|
||||
torch.tensor([0], dtype=cur_idxs.dtype, device=cur_idxs.device),
|
||||
cur_idxs + 1,
|
||||
torch.tensor([h * w], dtype=cur_idxs.dtype, device=cur_idxs.device),
|
||||
]
|
||||
)
|
||||
btw_idxs = cur_idxs[1:] - cur_idxs[:-1]
|
||||
counts = [] if tensor[i, 0] == 0 else [0]
|
||||
counts.extend(btw_idxs.detach().cpu().tolist())
|
||||
out.append({"size": [h, w], "counts": counts})
|
||||
return out
|
||||
|
||||
|
||||
def rle_to_mask(rle: Dict[str, Any]) -> np.ndarray:
|
||||
"""Compute a binary mask from an uncompressed RLE."""
|
||||
h, w = rle["size"]
|
||||
mask = np.empty(h * w, dtype=bool)
|
||||
idx = 0
|
||||
parity = False
|
||||
for count in rle["counts"]:
|
||||
mask[idx : idx + count] = parity
|
||||
idx += count
|
||||
parity ^= True
|
||||
mask = mask.reshape(w, h)
|
||||
return mask.transpose() # Put in C order
|
||||
|
||||
|
||||
def area_from_rle(rle: Dict[str, Any]) -> int:
|
||||
return sum(rle["counts"][1::2])
|
||||
|
||||
|
||||
def calculate_stability_score(
|
||||
masks: torch.Tensor, mask_threshold: float, threshold_offset: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Computes the stability score for a batch of masks. The stability
|
||||
score is the IoU between the binary masks obtained by thresholding
|
||||
the predicted mask logits at high and low values.
|
||||
"""
|
||||
# One mask is always contained inside the other.
|
||||
# Save memory by preventing unnecessary cast to torch.int64
|
||||
intersections = (
|
||||
(masks > (mask_threshold + threshold_offset))
|
||||
.sum(-1, dtype=torch.int16)
|
||||
.sum(-1, dtype=torch.int32)
|
||||
)
|
||||
unions = (
|
||||
(masks > (mask_threshold - threshold_offset))
|
||||
.sum(-1, dtype=torch.int16)
|
||||
.sum(-1, dtype=torch.int32)
|
||||
)
|
||||
return intersections / unions
|
||||
|
||||
|
||||
def build_point_grid(n_per_side: int) -> np.ndarray:
|
||||
"""Generates a 2D grid of points evenly spaced in [0,1]x[0,1]."""
|
||||
offset = 1 / (2 * n_per_side)
|
||||
points_one_side = np.linspace(offset, 1 - offset, n_per_side)
|
||||
points_x = np.tile(points_one_side[None, :], (n_per_side, 1))
|
||||
points_y = np.tile(points_one_side[:, None], (1, n_per_side))
|
||||
points = np.stack([points_x, points_y], axis=-1).reshape(-1, 2)
|
||||
return points
|
||||
|
||||
|
||||
def build_all_layer_point_grids(
|
||||
n_per_side: int, n_layers: int, scale_per_layer: int
|
||||
) -> List[np.ndarray]:
|
||||
"""Generates point grids for all crop layers."""
|
||||
points_by_layer = []
|
||||
for i in range(n_layers + 1):
|
||||
n_points = int(n_per_side / (scale_per_layer**i))
|
||||
points_by_layer.append(build_point_grid(n_points))
|
||||
return points_by_layer
|
||||
|
||||
|
||||
def generate_crop_boxes(
|
||||
im_size: Tuple[int, ...], n_layers: int, overlap_ratio: float
|
||||
) -> Tuple[List[List[int]], List[int]]:
|
||||
"""
|
||||
Generates a list of crop boxes of different sizes. Each layer
|
||||
has (2**i)**2 boxes for the ith layer.
|
||||
"""
|
||||
crop_boxes, layer_idxs = [], []
|
||||
im_h, im_w = im_size
|
||||
short_side = min(im_h, im_w)
|
||||
|
||||
# Original image
|
||||
crop_boxes.append([0, 0, im_w, im_h])
|
||||
layer_idxs.append(0)
|
||||
|
||||
def crop_len(orig_len, n_crops, overlap):
|
||||
return int(math.ceil((overlap * (n_crops - 1) + orig_len) / n_crops))
|
||||
|
||||
for i_layer in range(n_layers):
|
||||
n_crops_per_side = 2 ** (i_layer + 1)
|
||||
overlap = int(overlap_ratio * short_side * (2 / n_crops_per_side))
|
||||
|
||||
crop_w = crop_len(im_w, n_crops_per_side, overlap)
|
||||
crop_h = crop_len(im_h, n_crops_per_side, overlap)
|
||||
|
||||
crop_box_x0 = [int((crop_w - overlap) * i) for i in range(n_crops_per_side)]
|
||||
crop_box_y0 = [int((crop_h - overlap) * i) for i in range(n_crops_per_side)]
|
||||
|
||||
# Crops in XYWH format
|
||||
for x0, y0 in product(crop_box_x0, crop_box_y0):
|
||||
box = [x0, y0, min(x0 + crop_w, im_w), min(y0 + crop_h, im_h)]
|
||||
crop_boxes.append(box)
|
||||
layer_idxs.append(i_layer + 1)
|
||||
|
||||
return crop_boxes, layer_idxs
|
||||
|
||||
|
||||
def uncrop_boxes_xyxy(boxes: torch.Tensor, crop_box: List[int]) -> torch.Tensor:
|
||||
x0, y0, _, _ = crop_box
|
||||
offset = torch.tensor([[x0, y0, x0, y0]], device=boxes.device)
|
||||
# Check if boxes has a channel dimension
|
||||
if len(boxes.shape) == 3:
|
||||
offset = offset.unsqueeze(1)
|
||||
return boxes + offset
|
||||
|
||||
|
||||
def uncrop_points(points: torch.Tensor, crop_box: List[int]) -> torch.Tensor:
|
||||
x0, y0, _, _ = crop_box
|
||||
offset = torch.tensor([[x0, y0]], device=points.device)
|
||||
# Check if points has a channel dimension
|
||||
if len(points.shape) == 3:
|
||||
offset = offset.unsqueeze(1)
|
||||
return points + offset
|
||||
|
||||
|
||||
def uncrop_masks(
|
||||
masks: torch.Tensor, crop_box: List[int], orig_h: int, orig_w: int
|
||||
) -> torch.Tensor:
|
||||
x0, y0, x1, y1 = crop_box
|
||||
if x0 == 0 and y0 == 0 and x1 == orig_w and y1 == orig_h:
|
||||
return masks
|
||||
# Coordinate transform masks
|
||||
pad_x, pad_y = orig_w - (x1 - x0), orig_h - (y1 - y0)
|
||||
pad = (x0, pad_x - x0, y0, pad_y - y0)
|
||||
return torch.nn.functional.pad(masks, pad, value=0)
|
||||
|
||||
|
||||
def remove_small_regions(
|
||||
mask: np.ndarray, area_thresh: float, mode: str
|
||||
) -> Tuple[np.ndarray, bool]:
|
||||
"""
|
||||
Removes small disconnected regions and holes in a mask. Returns the
|
||||
mask and an indicator of if the mask has been modified.
|
||||
"""
|
||||
import cv2 # type: ignore
|
||||
|
||||
assert mode in ["holes", "islands"]
|
||||
correct_holes = mode == "holes"
|
||||
working_mask = (correct_holes ^ mask).astype(np.uint8)
|
||||
n_labels, regions, stats, _ = cv2.connectedComponentsWithStats(working_mask, 8)
|
||||
sizes = stats[:, -1][1:] # Row 0 is background label
|
||||
small_regions = [i + 1 for i, s in enumerate(sizes) if s < area_thresh]
|
||||
if len(small_regions) == 0:
|
||||
return mask, False
|
||||
fill_labels = [0] + small_regions
|
||||
if not correct_holes:
|
||||
fill_labels = [i for i in range(n_labels) if i not in fill_labels]
|
||||
# If every region is below threshold, keep largest
|
||||
if len(fill_labels) == 0:
|
||||
fill_labels = [int(np.argmax(sizes)) + 1]
|
||||
mask = np.isin(regions, fill_labels)
|
||||
return mask, True
|
||||
|
||||
|
||||
def coco_encode_rle(uncompressed_rle: Dict[str, Any]) -> Dict[str, Any]:
|
||||
from pycocotools import mask as mask_utils # type: ignore
|
||||
|
||||
h, w = uncompressed_rle["size"]
|
||||
rle = mask_utils.frPyObjects(uncompressed_rle, h, w)
|
||||
rle["counts"] = rle["counts"].decode("utf-8") # Necessary to serialize with json
|
||||
return rle
|
||||
|
||||
|
||||
def batched_mask_to_box(masks: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Calculates boxes in XYXY format around masks. Return [0,0,0,0] for
|
||||
an empty mask. For input shape C1xC2x...xHxW, the output shape is C1xC2x...x4.
|
||||
"""
|
||||
# torch.max below raises an error on empty inputs, just skip in this case
|
||||
if torch.numel(masks) == 0:
|
||||
return torch.zeros(*masks.shape[:-2], 4, device=masks.device)
|
||||
|
||||
# Normalize shape to CxHxW
|
||||
shape = masks.shape
|
||||
h, w = shape[-2:]
|
||||
if len(shape) > 2:
|
||||
masks = masks.flatten(0, -3)
|
||||
else:
|
||||
masks = masks.unsqueeze(0)
|
||||
|
||||
# Get top and bottom edges
|
||||
in_height, _ = torch.max(masks, dim=-1)
|
||||
in_height_coords = in_height * torch.arange(h, device=in_height.device)[None, :]
|
||||
bottom_edges, _ = torch.max(in_height_coords, dim=-1)
|
||||
in_height_coords = in_height_coords + h * (~in_height)
|
||||
top_edges, _ = torch.min(in_height_coords, dim=-1)
|
||||
|
||||
# Get left and right edges
|
||||
in_width, _ = torch.max(masks, dim=-2)
|
||||
in_width_coords = in_width * torch.arange(w, device=in_width.device)[None, :]
|
||||
right_edges, _ = torch.max(in_width_coords, dim=-1)
|
||||
in_width_coords = in_width_coords + w * (~in_width)
|
||||
left_edges, _ = torch.min(in_width_coords, dim=-1)
|
||||
|
||||
# If the mask is empty the right edge will be to the left of the left edge.
|
||||
# Replace these boxes with [0, 0, 0, 0]
|
||||
empty_filter = (right_edges < left_edges) | (bottom_edges < top_edges)
|
||||
out = torch.stack([left_edges, top_edges, right_edges, bottom_edges], dim=-1)
|
||||
out = out * (~empty_filter).unsqueeze(-1)
|
||||
|
||||
# Return to original shape
|
||||
if len(shape) > 2:
|
||||
out = out.reshape(*shape[:-2], 4)
|
||||
else:
|
||||
out = out[0]
|
||||
|
||||
return out
|
||||
@@ -0,0 +1,238 @@
|
||||
# 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
|
||||
|
||||
|
||||
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 model.segment_anything_2.sam2 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] boxes, 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):
|
||||
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
|
||||
|
||||
# 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.cuda(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,
|
||||
):
|
||||
"""
|
||||
Load the video frames from a directory of JPEG files ("<frame_index>.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")
|
||||
|
||||
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
|
||||
)
|
||||
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.cuda()
|
||||
img_mean = img_mean.cuda()
|
||||
img_std = img_std.cuda()
|
||||
# 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"
|
||||
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)
|
||||
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}
|
||||
@@ -0,0 +1,99 @@
|
||||
# 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 torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torchvision.transforms import Normalize, Resize, ToTensor
|
||||
|
||||
|
||||
class SAM2Transforms(nn.Module):
|
||||
def __init__(
|
||||
self, resolution, mask_threshold, max_hole_area=0.0, max_sprinkle_area=0.0
|
||||
):
|
||||
"""
|
||||
Transforms for SAM2.
|
||||
"""
|
||||
super().__init__()
|
||||
self.resolution = resolution
|
||||
self.mask_threshold = mask_threshold
|
||||
self.max_hole_area = max_hole_area
|
||||
self.max_sprinkle_area = max_sprinkle_area
|
||||
self.mean = [0.485, 0.456, 0.406]
|
||||
self.std = [0.229, 0.224, 0.225]
|
||||
self.to_tensor = ToTensor()
|
||||
self.transforms = torch.jit.script(
|
||||
nn.Sequential(
|
||||
Resize((self.resolution, self.resolution)),
|
||||
Normalize(self.mean, self.std),
|
||||
)
|
||||
)
|
||||
|
||||
def __call__(self, x):
|
||||
x = self.to_tensor(x)
|
||||
return self.transforms(x)
|
||||
|
||||
def forward_batch(self, img_list):
|
||||
img_batch = [self.transforms(self.to_tensor(img)) for img in img_list]
|
||||
img_batch = torch.stack(img_batch, dim=0)
|
||||
return img_batch
|
||||
|
||||
def transform_coords(
|
||||
self, coords: torch.Tensor, normalize=False, orig_hw=None
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Expects a torch tensor with length 2 in the last dimension. The coordinates can be in absolute image or normalized coordinates,
|
||||
If the coords are in absolute image coordinates, normalize should be set to True and original image size is required.
|
||||
|
||||
Returns
|
||||
Un-normalized coordinates in the range of [0, 1] which is expected by the SAM2 model.
|
||||
"""
|
||||
if normalize:
|
||||
assert orig_hw is not None
|
||||
h, w = orig_hw
|
||||
coords = coords.clone()
|
||||
coords[..., 0] = coords[..., 0] / w
|
||||
coords[..., 1] = coords[..., 1] / h
|
||||
|
||||
coords = coords * self.resolution # unnormalize coords
|
||||
return coords
|
||||
|
||||
def transform_boxes(
|
||||
self, boxes: torch.Tensor, normalize=False, orig_hw=None
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Expects a tensor of shape Bx4. The coordinates can be in absolute image or normalized coordinates,
|
||||
if the coords are in absolute image coordinates, normalize should be set to True and original image size is required.
|
||||
"""
|
||||
boxes = self.transform_coords(boxes.reshape(-1, 2, 2), normalize, orig_hw)
|
||||
return boxes
|
||||
|
||||
def postprocess_masks(self, masks: torch.Tensor, orig_hw) -> torch.Tensor:
|
||||
"""
|
||||
Perform PostProcessing on output masks.
|
||||
"""
|
||||
from model.segment_anything_2.sam2.utils.misc import get_connected_components
|
||||
|
||||
masks = masks.float()
|
||||
if self.max_hole_area > 0:
|
||||
# Holes are those connected components in background with area <= self.fill_hole_area
|
||||
# (background regions are those with mask scores <= self.mask_threshold)
|
||||
mask_flat = masks.flatten(0, 1).unsqueeze(1) # flatten as 1-channel image
|
||||
labels, areas = get_connected_components(mask_flat <= self.mask_threshold)
|
||||
is_hole = (labels > 0) & (areas <= self.max_hole_area)
|
||||
is_hole = is_hole.reshape_as(masks)
|
||||
# We fill holes with a small positive mask score (10.0) to change them to foreground.
|
||||
masks = torch.where(is_hole, self.mask_threshold + 10.0, masks)
|
||||
|
||||
if self.max_sprinkle_area > 0:
|
||||
labels, areas = get_connected_components(mask_flat > self.mask_threshold)
|
||||
is_hole = (labels > 0) & (areas <= self.max_sprinkle_area)
|
||||
is_hole = is_hole.reshape_as(masks)
|
||||
# We fill holes with negative mask score (-10.0) to change them to background.
|
||||
masks = torch.where(is_hole, self.mask_threshold - 10.0, masks)
|
||||
|
||||
masks = F.interpolate(masks, orig_hw, mode="bilinear", align_corners=False)
|
||||
return masks
|
||||
@@ -0,0 +1,5 @@
|
||||
# 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.
|
||||
@@ -0,0 +1,113 @@
|
||||
# @package _global_
|
||||
|
||||
# Model
|
||||
model:
|
||||
_target_: model.segment_anything_2.sam2.modeling.sam2_base.SAM2Base
|
||||
image_encoder:
|
||||
_target_: model.segment_anything_2.sam2.modeling.backbones.image_encoder.ImageEncoder
|
||||
scalp: 1
|
||||
trunk:
|
||||
_target_: model.segment_anything_2.sam2.modeling.backbones.hieradet.Hiera
|
||||
embed_dim: 112
|
||||
num_heads: 2
|
||||
neck:
|
||||
_target_: model.segment_anything_2.sam2.modeling.backbones.image_encoder.FpnNeck
|
||||
position_encoding:
|
||||
_target_: model.segment_anything_2.sam2.modeling.position_encoding.PositionEmbeddingSine
|
||||
num_pos_feats: 256
|
||||
normalize: true
|
||||
scale: null
|
||||
temperature: 10000
|
||||
d_model: 256
|
||||
backbone_channel_list: [896, 448, 224, 112]
|
||||
fpn_top_down_levels: [2, 3] # output level 0 and 1 directly use the backbone features
|
||||
fpn_interp_model: nearest
|
||||
|
||||
memory_attention:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_attention.MemoryAttention
|
||||
d_model: 256
|
||||
pos_enc_at_input: true
|
||||
layer:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_attention.MemoryAttentionLayer
|
||||
activation: relu
|
||||
dim_feedforward: 2048
|
||||
dropout: 0.1
|
||||
pos_enc_at_attn: false
|
||||
self_attention:
|
||||
_target_: model.segment_anything_2.sam2.modeling.sam.transformer.RoPEAttention
|
||||
rope_theta: 10000.0
|
||||
feat_sizes: [32, 32]
|
||||
embedding_dim: 256
|
||||
num_heads: 1
|
||||
downsample_rate: 1
|
||||
dropout: 0.1
|
||||
d_model: 256
|
||||
pos_enc_at_cross_attn_keys: true
|
||||
pos_enc_at_cross_attn_queries: false
|
||||
cross_attention:
|
||||
_target_: model.segment_anything_2.sam2.modeling.sam.transformer.RoPEAttention
|
||||
rope_theta: 10000.0
|
||||
feat_sizes: [32, 32]
|
||||
rope_k_repeat: True
|
||||
embedding_dim: 256
|
||||
num_heads: 1
|
||||
downsample_rate: 1
|
||||
dropout: 0.1
|
||||
kv_in_dim: 64
|
||||
num_layers: 4
|
||||
|
||||
memory_encoder:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_encoder.MemoryEncoder
|
||||
out_dim: 64
|
||||
position_encoding:
|
||||
_target_: model.segment_anything_2.sam2.modeling.position_encoding.PositionEmbeddingSine
|
||||
num_pos_feats: 64
|
||||
normalize: true
|
||||
scale: null
|
||||
temperature: 10000
|
||||
mask_downsampler:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_encoder.MaskDownSampler
|
||||
kernel_size: 3
|
||||
stride: 2
|
||||
padding: 1
|
||||
fuser:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_encoder.Fuser
|
||||
layer:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_encoder.CXBlock
|
||||
dim: 256
|
||||
kernel_size: 7
|
||||
padding: 3
|
||||
layer_scale_init_value: 1e-6
|
||||
use_dwconv: True # depth-wise convs
|
||||
num_layers: 2
|
||||
|
||||
num_maskmem: 7
|
||||
image_size: 1024
|
||||
# apply scaled sigmoid on mask logits for memory encoder, and directly feed input mask as output mask
|
||||
sigmoid_scale_for_mem_enc: 20.0
|
||||
sigmoid_bias_for_mem_enc: -10.0
|
||||
use_mask_input_as_output_without_sam: true
|
||||
# Memory
|
||||
directly_add_no_mem_embed: true
|
||||
# use high-resolution feature map in the SAM mask decoder
|
||||
use_high_res_features_in_sam: true
|
||||
# output 3 masks on the first click on initial conditioning frames
|
||||
multimask_output_in_sam: true
|
||||
# SAM heads
|
||||
iou_prediction_use_sigmoid: True
|
||||
# cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
|
||||
use_obj_ptrs_in_encoder: true
|
||||
add_tpos_enc_to_obj_ptrs: false
|
||||
only_obj_ptrs_in_the_past_for_eval: true
|
||||
# object occlusion prediction
|
||||
pred_obj_scores: true
|
||||
pred_obj_scores_mlp: true
|
||||
fixed_no_obj_ptr: true
|
||||
# multimask tracking settings
|
||||
multimask_output_for_tracking: true
|
||||
use_multimask_token_for_obj_ptr: true
|
||||
multimask_min_pt_num: 0
|
||||
multimask_max_pt_num: 1
|
||||
use_mlp_for_obj_ptr_proj: true
|
||||
# Compilation flag
|
||||
compile_image_encoder: False
|
||||
@@ -0,0 +1,117 @@
|
||||
# @package _global_
|
||||
|
||||
# Model
|
||||
model:
|
||||
_target_: model.segment_anything_2.sam2.modeling.sam2_base.SAM2Base
|
||||
image_encoder:
|
||||
_target_: model.segment_anything_2.sam2.modeling.backbones.image_encoder.ImageEncoder
|
||||
scalp: 1
|
||||
trunk:
|
||||
_target_: model.segment_anything_2.sam2.modeling.backbones.hieradet.Hiera
|
||||
embed_dim: 144
|
||||
num_heads: 2
|
||||
stages: [2, 6, 36, 4]
|
||||
global_att_blocks: [23, 33, 43]
|
||||
window_pos_embed_bkg_spatial_size: [7, 7]
|
||||
window_spec: [8, 4, 16, 8]
|
||||
neck:
|
||||
_target_: model.segment_anything_2.sam2.modeling.backbones.image_encoder.FpnNeck
|
||||
position_encoding:
|
||||
_target_: model.segment_anything_2.sam2.modeling.position_encoding.PositionEmbeddingSine
|
||||
num_pos_feats: 256
|
||||
normalize: true
|
||||
scale: null
|
||||
temperature: 10000
|
||||
d_model: 256
|
||||
backbone_channel_list: [1152, 576, 288, 144]
|
||||
fpn_top_down_levels: [2, 3] # output level 0 and 1 directly use the backbone features
|
||||
fpn_interp_model: nearest
|
||||
|
||||
memory_attention:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_attention.MemoryAttention
|
||||
d_model: 256
|
||||
pos_enc_at_input: true
|
||||
layer:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_attention.MemoryAttentionLayer
|
||||
activation: relu
|
||||
dim_feedforward: 2048
|
||||
dropout: 0.1
|
||||
pos_enc_at_attn: false
|
||||
self_attention:
|
||||
_target_: model.segment_anything_2.sam2.modeling.sam.transformer.RoPEAttention
|
||||
rope_theta: 10000.0
|
||||
feat_sizes: [32, 32]
|
||||
embedding_dim: 256
|
||||
num_heads: 1
|
||||
downsample_rate: 1
|
||||
dropout: 0.1
|
||||
d_model: 256
|
||||
pos_enc_at_cross_attn_keys: true
|
||||
pos_enc_at_cross_attn_queries: false
|
||||
cross_attention:
|
||||
_target_: model.segment_anything_2.sam2.modeling.sam.transformer.RoPEAttention
|
||||
rope_theta: 10000.0
|
||||
feat_sizes: [32, 32]
|
||||
rope_k_repeat: True
|
||||
embedding_dim: 256
|
||||
num_heads: 1
|
||||
downsample_rate: 1
|
||||
dropout: 0.1
|
||||
kv_in_dim: 64
|
||||
num_layers: 4
|
||||
|
||||
memory_encoder:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_encoder.MemoryEncoder
|
||||
out_dim: 64
|
||||
position_encoding:
|
||||
_target_: model.segment_anything_2.sam2.modeling.position_encoding.PositionEmbeddingSine
|
||||
num_pos_feats: 64
|
||||
normalize: true
|
||||
scale: null
|
||||
temperature: 10000
|
||||
mask_downsampler:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_encoder.MaskDownSampler
|
||||
kernel_size: 3
|
||||
stride: 2
|
||||
padding: 1
|
||||
fuser:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_encoder.Fuser
|
||||
layer:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_encoder.CXBlock
|
||||
dim: 256
|
||||
kernel_size: 7
|
||||
padding: 3
|
||||
layer_scale_init_value: 1e-6
|
||||
use_dwconv: True # depth-wise convs
|
||||
num_layers: 2
|
||||
|
||||
num_maskmem: 7
|
||||
image_size: 1024
|
||||
# apply scaled sigmoid on mask logits for memory encoder, and directly feed input mask as output mask
|
||||
sigmoid_scale_for_mem_enc: 20.0
|
||||
sigmoid_bias_for_mem_enc: -10.0
|
||||
use_mask_input_as_output_without_sam: true
|
||||
# Memory
|
||||
directly_add_no_mem_embed: true
|
||||
# use high-resolution feature map in the SAM mask decoder
|
||||
use_high_res_features_in_sam: true
|
||||
# output 3 masks on the first click on initial conditioning frames
|
||||
multimask_output_in_sam: true
|
||||
# SAM heads
|
||||
iou_prediction_use_sigmoid: True
|
||||
# cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
|
||||
use_obj_ptrs_in_encoder: true
|
||||
add_tpos_enc_to_obj_ptrs: false
|
||||
only_obj_ptrs_in_the_past_for_eval: true
|
||||
# object occlusion prediction
|
||||
pred_obj_scores: true
|
||||
pred_obj_scores_mlp: true
|
||||
fixed_no_obj_ptr: true
|
||||
# multimask tracking settings
|
||||
multimask_output_for_tracking: true
|
||||
use_multimask_token_for_obj_ptr: true
|
||||
multimask_min_pt_num: 0
|
||||
multimask_max_pt_num: 1
|
||||
use_mlp_for_obj_ptr_proj: true
|
||||
# Compilation flag
|
||||
compile_image_encoder: False
|
||||
@@ -0,0 +1,116 @@
|
||||
# @package _global_
|
||||
|
||||
# Model
|
||||
model:
|
||||
_target_: model.segment_anything_2.sam2.modeling.sam2_base.SAM2Base
|
||||
image_encoder:
|
||||
_target_: model.segment_anything_2.sam2.modeling.backbones.image_encoder.ImageEncoder
|
||||
scalp: 1
|
||||
trunk:
|
||||
_target_: model.segment_anything_2.sam2.modeling.backbones.hieradet.Hiera
|
||||
embed_dim: 96
|
||||
num_heads: 1
|
||||
stages: [1, 2, 11, 2]
|
||||
global_att_blocks: [7, 10, 13]
|
||||
window_pos_embed_bkg_spatial_size: [7, 7]
|
||||
neck:
|
||||
_target_: model.segment_anything_2.sam2.modeling.backbones.image_encoder.FpnNeck
|
||||
position_encoding:
|
||||
_target_: model.segment_anything_2.sam2.modeling.position_encoding.PositionEmbeddingSine
|
||||
num_pos_feats: 256
|
||||
normalize: true
|
||||
scale: null
|
||||
temperature: 10000
|
||||
d_model: 256
|
||||
backbone_channel_list: [768, 384, 192, 96]
|
||||
fpn_top_down_levels: [2, 3] # output level 0 and 1 directly use the backbone features
|
||||
fpn_interp_model: nearest
|
||||
|
||||
memory_attention:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_attention.MemoryAttention
|
||||
d_model: 256
|
||||
pos_enc_at_input: true
|
||||
layer:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_attention.MemoryAttentionLayer
|
||||
activation: relu
|
||||
dim_feedforward: 2048
|
||||
dropout: 0.1
|
||||
pos_enc_at_attn: false
|
||||
self_attention:
|
||||
_target_: model.segment_anything_2.sam2.modeling.sam.transformer.RoPEAttention
|
||||
rope_theta: 10000.0
|
||||
feat_sizes: [32, 32]
|
||||
embedding_dim: 256
|
||||
num_heads: 1
|
||||
downsample_rate: 1
|
||||
dropout: 0.1
|
||||
d_model: 256
|
||||
pos_enc_at_cross_attn_keys: true
|
||||
pos_enc_at_cross_attn_queries: false
|
||||
cross_attention:
|
||||
_target_: model.segment_anything_2.sam2.modeling.sam.transformer.RoPEAttention
|
||||
rope_theta: 10000.0
|
||||
feat_sizes: [32, 32]
|
||||
rope_k_repeat: True
|
||||
embedding_dim: 256
|
||||
num_heads: 1
|
||||
downsample_rate: 1
|
||||
dropout: 0.1
|
||||
kv_in_dim: 64
|
||||
num_layers: 4
|
||||
|
||||
memory_encoder:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_encoder.MemoryEncoder
|
||||
out_dim: 64
|
||||
position_encoding:
|
||||
_target_: model.segment_anything_2.sam2.modeling.position_encoding.PositionEmbeddingSine
|
||||
num_pos_feats: 64
|
||||
normalize: true
|
||||
scale: null
|
||||
temperature: 10000
|
||||
mask_downsampler:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_encoder.MaskDownSampler
|
||||
kernel_size: 3
|
||||
stride: 2
|
||||
padding: 1
|
||||
fuser:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_encoder.Fuser
|
||||
layer:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_encoder.CXBlock
|
||||
dim: 256
|
||||
kernel_size: 7
|
||||
padding: 3
|
||||
layer_scale_init_value: 1e-6
|
||||
use_dwconv: True # depth-wise convs
|
||||
num_layers: 2
|
||||
|
||||
num_maskmem: 7
|
||||
image_size: 1024
|
||||
# apply scaled sigmoid on mask logits for memory encoder, and directly feed input mask as output mask
|
||||
sigmoid_scale_for_mem_enc: 20.0
|
||||
sigmoid_bias_for_mem_enc: -10.0
|
||||
use_mask_input_as_output_without_sam: true
|
||||
# Memory
|
||||
directly_add_no_mem_embed: true
|
||||
# use high-resolution feature map in the SAM mask decoder
|
||||
use_high_res_features_in_sam: true
|
||||
# output 3 masks on the first click on initial conditioning frames
|
||||
multimask_output_in_sam: true
|
||||
# SAM heads
|
||||
iou_prediction_use_sigmoid: True
|
||||
# cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
|
||||
use_obj_ptrs_in_encoder: true
|
||||
add_tpos_enc_to_obj_ptrs: false
|
||||
only_obj_ptrs_in_the_past_for_eval: true
|
||||
# object occlusion prediction
|
||||
pred_obj_scores: true
|
||||
pred_obj_scores_mlp: true
|
||||
fixed_no_obj_ptr: true
|
||||
# multimask tracking settings
|
||||
multimask_output_for_tracking: true
|
||||
use_multimask_token_for_obj_ptr: true
|
||||
multimask_min_pt_num: 0
|
||||
multimask_max_pt_num: 1
|
||||
use_mlp_for_obj_ptr_proj: true
|
||||
# Compilation flag
|
||||
compile_image_encoder: False
|
||||
@@ -0,0 +1,118 @@
|
||||
# @package _global_
|
||||
|
||||
# Model
|
||||
model:
|
||||
_target_: model.segment_anything_2.sam2.modeling.sam2_base.SAM2Base
|
||||
image_encoder:
|
||||
_target_: model.segment_anything_2.sam2.modeling.backbones.image_encoder.ImageEncoder
|
||||
scalp: 1
|
||||
trunk:
|
||||
_target_: model.segment_anything_2.sam2.modeling.backbones.hieradet.Hiera
|
||||
embed_dim: 96
|
||||
num_heads: 1
|
||||
stages: [1, 2, 7, 2]
|
||||
global_att_blocks: [5, 7, 9]
|
||||
window_pos_embed_bkg_spatial_size: [7, 7]
|
||||
neck:
|
||||
_target_: model.segment_anything_2.sam2.modeling.backbones.image_encoder.FpnNeck
|
||||
position_encoding:
|
||||
_target_: model.segment_anything_2.sam2.modeling.position_encoding.PositionEmbeddingSine
|
||||
num_pos_feats: 256
|
||||
normalize: true
|
||||
scale: null
|
||||
temperature: 10000
|
||||
d_model: 256
|
||||
backbone_channel_list: [768, 384, 192, 96]
|
||||
fpn_top_down_levels: [2, 3] # output level 0 and 1 directly use the backbone features
|
||||
fpn_interp_model: nearest
|
||||
|
||||
memory_attention:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_attention.MemoryAttention
|
||||
d_model: 256
|
||||
pos_enc_at_input: true
|
||||
layer:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_attention.MemoryAttentionLayer
|
||||
activation: relu
|
||||
dim_feedforward: 2048
|
||||
dropout: 0.1
|
||||
pos_enc_at_attn: false
|
||||
self_attention:
|
||||
_target_: model.segment_anything_2.sam2.modeling.sam.transformer.RoPEAttention
|
||||
rope_theta: 10000.0
|
||||
feat_sizes: [32, 32]
|
||||
embedding_dim: 256
|
||||
num_heads: 1
|
||||
downsample_rate: 1
|
||||
dropout: 0.1
|
||||
d_model: 256
|
||||
pos_enc_at_cross_attn_keys: true
|
||||
pos_enc_at_cross_attn_queries: false
|
||||
cross_attention:
|
||||
_target_: model.segment_anything_2.sam2.modeling.sam.transformer.RoPEAttention
|
||||
rope_theta: 10000.0
|
||||
feat_sizes: [32, 32]
|
||||
rope_k_repeat: True
|
||||
embedding_dim: 256
|
||||
num_heads: 1
|
||||
downsample_rate: 1
|
||||
dropout: 0.1
|
||||
kv_in_dim: 64
|
||||
num_layers: 4
|
||||
|
||||
memory_encoder:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_encoder.MemoryEncoder
|
||||
out_dim: 64
|
||||
position_encoding:
|
||||
_target_: model.segment_anything_2.sam2.modeling.position_encoding.PositionEmbeddingSine
|
||||
num_pos_feats: 64
|
||||
normalize: true
|
||||
scale: null
|
||||
temperature: 10000
|
||||
mask_downsampler:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_encoder.MaskDownSampler
|
||||
kernel_size: 3
|
||||
stride: 2
|
||||
padding: 1
|
||||
fuser:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_encoder.Fuser
|
||||
layer:
|
||||
_target_: model.segment_anything_2.sam2.modeling.memory_encoder.CXBlock
|
||||
dim: 256
|
||||
kernel_size: 7
|
||||
padding: 3
|
||||
layer_scale_init_value: 1e-6
|
||||
use_dwconv: True # depth-wise convs
|
||||
num_layers: 2
|
||||
|
||||
num_maskmem: 7
|
||||
image_size: 1024
|
||||
# apply scaled sigmoid on mask logits for memory encoder, and directly feed input mask as output mask
|
||||
# SAM decoder
|
||||
sigmoid_scale_for_mem_enc: 20.0
|
||||
sigmoid_bias_for_mem_enc: -10.0
|
||||
use_mask_input_as_output_without_sam: true
|
||||
# Memory
|
||||
directly_add_no_mem_embed: true
|
||||
# use high-resolution feature map in the SAM mask decoder
|
||||
use_high_res_features_in_sam: true
|
||||
# output 3 masks on the first click on initial conditioning frames
|
||||
multimask_output_in_sam: true
|
||||
# SAM heads
|
||||
iou_prediction_use_sigmoid: True
|
||||
# cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder
|
||||
use_obj_ptrs_in_encoder: true
|
||||
add_tpos_enc_to_obj_ptrs: false
|
||||
only_obj_ptrs_in_the_past_for_eval: true
|
||||
# object occlusion prediction
|
||||
pred_obj_scores: true
|
||||
pred_obj_scores_mlp: true
|
||||
fixed_no_obj_ptr: true
|
||||
# multimask tracking settings
|
||||
multimask_output_for_tracking: true
|
||||
use_multimask_token_for_obj_ptr: true
|
||||
multimask_min_pt_num: 0
|
||||
multimask_max_pt_num: 1
|
||||
use_mlp_for_obj_ptr_proj: true
|
||||
# Compilation flag
|
||||
# HieraT does not currently support compilation, should always be set to False
|
||||
compile_image_encoder: False
|
||||
@@ -0,0 +1,29 @@
|
||||
# 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.
|
||||
|
||||
from setuptools import find_packages, setup
|
||||
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
||||
|
||||
def get_extensions():
|
||||
srcs = ["sam2/csrc/connected_components.cu"]
|
||||
compile_args = {
|
||||
"cxx": [],
|
||||
"nvcc": [
|
||||
"-DCUDA_HAS_FP16=1",
|
||||
"-D__CUDA_NO_HALF_OPERATORS__",
|
||||
"-D__CUDA_NO_HALF_CONVERSIONS__",
|
||||
"-D__CUDA_NO_HALF2_OPERATORS__",
|
||||
],
|
||||
}
|
||||
ext_modules = [CUDAExtension("sam2._C", srcs, extra_compile_args=compile_args)]
|
||||
return ext_modules
|
||||
|
||||
|
||||
# Setup configuration
|
||||
setup(
|
||||
ext_modules=get_extensions(),
|
||||
cmdclass={"build_ext": BuildExtension.with_options(no_python_abi_suffix=True)},
|
||||
)
|
||||
@@ -0,0 +1,191 @@
|
||||
# [(BEiT-3) Image as a Foreign Language: BEiT Pretraining for Vision and Vision-Language Tasks](https://arxiv.org/abs/2208.10442)
|
||||
|
||||
Official PyTorch implementation and pretrained models of BEiT-3.
|
||||
|
||||
The code and pretrained models of **BEiT** can be found at [here](https://github.com/microsoft/unilm/tree/master/beit).
|
||||
|
||||
The code and pretrained models of **BEiT v2** can be found at [here](https://github.com/microsoft/unilm/tree/master/beit2).
|
||||
|
||||
- March, 2023: release [the code and pretrained models of **BEiT-3**](https://github.com/microsoft/unilm/tree/master/beit3)
|
||||
- March, 2023: [**BEiT-3**](https://arxiv.org/abs/2208.10442) was accepted by **CVPR 2023**.
|
||||
- Sept 2022: release [the code and pretrained models of **BEiT v2**](https://github.com/microsoft/unilm/tree/master/beit2)
|
||||
- Aug 2022: release preprint [Image as a Foreign Language: BEiT Pretraining for All Vision and Vision-Language Tasks](https://arxiv.org/abs/2208.10442)
|
||||
- Aug 2022: release preprint [BEiT v2: Masked Image Modeling with Vector-Quantized Visual Tokenizers](https://arxiv.org/abs/2208.06366)
|
||||
- June 2022: release preprint [VL-BEiT: Generative Vision-Language Pretraining](https://arxiv.org/abs/2206.01127)
|
||||
- March, 2022: add [linear probe examples](https://github.com/microsoft/unilm/blob/master/beit/get_started_for_image_classification.md#example-linear-probe-on-imagenet)
|
||||
- January, 2022: [**BEiT**](https://openreview.net/forum?id=p-BhZSz59o4) was accepted by **ICLR 2022 as Oral presentation** (54 out of 3391).
|
||||
- August 2021: [**BEiT**](https://huggingface.co/transformers/master/model_doc/beit.html) is on [HuggingFace](https://github.com/huggingface/transformers)
|
||||
- July 2021: BEiT-large achieves **[state-of-the-art results on ADE20K](https://paperswithcode.com/sota/semantic-segmentation-on-ade20k) (a big jump to 57.0 mIoU) for semantic segmentation**.
|
||||
- July 2021: BEiT-large achieves **state-of-the-art ImageNet top-1 accuracy (88.6%) under the setting without extra data other than ImageNet-22k**.
|
||||
- July 2021: release [the code and pretrained models of **BEiT**](https://github.com/microsoft/unilm/tree/master/beit)
|
||||
- June 2021: release preprint [BEiT: BERT Pre-Training of Image Transformers](https://arxiv.org/abs/2106.08254)
|
||||
|
||||
## Pretrained models
|
||||
|
||||
We provide BEiT-3 weights pretrained on monomodal and multimodal data. Our large-size model outperforms previous large-size models across various vision-language and vision downstream tasks. The models were pretrained with 224x224 resolution.
|
||||
|
||||
### Tips
|
||||
- For vision-language tasks that require deep fusion, we recommend using `BEiT3-base` and `BEiT3-large`.
|
||||
- For image-text retrieval or vision tasks, using `BEiT3-base-itc` and `BEiT3-large-itc` usually achieve better performance.
|
||||
|
||||
### Download Checkpoints
|
||||
|
||||
1. Models pretrained on ImageNet-21k images, 160 GB text documents, and web-scale image-text pairs (collected from [LAION-400M](https://laion.ai/blog/laion-400-open-dataset/), [English LAION-2B](https://laion.ai/blog/laion-5b/), [COYO-700M](https://github.com/kakaobrain/coyo-dataset), and CC15M).
|
||||
- [`BEiT3-base`](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_224.pth): #layer=12; hidden=768; FFN factor=4x; #head=12; patch=16x16; #parameters: 276M
|
||||
- [`BEiT3-large`](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_224.pth): #layer=24; hidden=1024; FFN factor=4x; #head=16; patch=16x16; #parameters: 746M
|
||||
|
||||
2. Perform image-text contrastive intermediate tuning on `BEiT3-base` and `BEiT3-large`.
|
||||
- [`BEiT3-base-itc`](https://github.com/addf400/files/releases/download/beit3/beit3_base_itc_patch16_224.pth): #layer=12; hidden=768; FFN factor=4x; #head=12; patch=16x16; #parameters: 222M
|
||||
- [`BEiT3-large-itc`](https://github.com/addf400/files/releases/download/beit3/beit3_large_itc_patch16_224.pth): #layer=24; hidden=1024; FFN factor=4x; #head=16; patch=16x16; #parameters: 674M
|
||||
|
||||
3. Add indomain image-text pairs (COCO and VG) to continue training `BEiT3-base` and `BEiT3-large` using masked data modeling. The indomain models achieve better performance on VQAv2 and NLVR2 tasks.
|
||||
- [`BEiT3-base-indomain`](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_224.pth): #layer=12; hidden=768; FFN factor=4x; #head=12; patch=16x16; #parameters: 276M
|
||||
- [`BEiT3-large-indomain`](https://github.com/addf400/files/releases/download/beit3/beit3_large_indomain_patch16_224.pth): #layer=24; hidden=1024; FFN factor=4x; #head=16; patch=16x16; #parameters: 746M
|
||||
|
||||
### Text Tokenizer
|
||||
|
||||
[beit3.spm](https://github.com/addf400/files/releases/download/beit3/beit3.spm) is the sentencepiece model used for tokenizing texts.
|
||||
```
|
||||
from transformers import XLMRobertaTokenizer
|
||||
tokenizer = XLMRobertaTokenizer("/your_beit3_model_path/beit3.spm")
|
||||
```
|
||||
|
||||
### Architecture
|
||||
|
||||
We use [Magneto](https://arxiv.org/abs/2210.06423) with decoupled Multiway Transformer as the backbone architecture. Magneto can have better training stability and obtain better performance across modalities (such as vision, and language). The implementation is based on the [torchscale](https://github.com/microsoft/torchscale/blob/main/torchscale/model/BEiT3.py) package.
|
||||
|
||||
|
||||
## Setup
|
||||
|
||||
```
|
||||
alias=`whoami | cut -d'.' -f2`; docker run -it --rm --runtime=nvidia --ipc=host --privileged -v /home/${alias}:/home/${alias} pytorch/pytorch:1.8.1-cuda11.1-cudnn8-devel bash
|
||||
```
|
||||
|
||||
Clone the repo and install required packages:
|
||||
```
|
||||
git clone https://github.com/microsoft/unilm.git
|
||||
cd unilm/beit3
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
|
||||
## Fine-tuning on ImageNet-1k (Image Classification)
|
||||
|
||||
The detailed instructions can be found at [`get_started_for_image_classification.md`](get_started/get_started_for_image_classification.md). We only use vision-related parameters for image classification fine-tuning.
|
||||
|
||||
| initialized checkpoint | resolution | acc@1 | acc@5 | #params | weight |
|
||||
|:----------------------------------------|:----------:|:-----:|:-----:|:-------:|-------------------|
|
||||
| [beit3_base_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_224.pth) | 224x224 | 85.4 | 97.6 | 87M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_224_in1k.pth) |
|
||||
| [beit3_base_indomain_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_224.pth) | 224x224 | 85.4 | 97.6 | 87M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_224_in1k.pth) |
|
||||
| [beit3_large_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_224.pth) | 224x224 | 87.6 | 98.3 | 305M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_224_in1k.pth) |
|
||||
| [beit3_large_indomain_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_indomain_patch16_224.pth) | 224x224 | 87.5 | 98.3 | 305M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_indomain_patch16_224_in1k.pth) |
|
||||
|
||||
|
||||
## Fine-tuning on VQAv2 (Visual Question Answering)
|
||||
|
||||
The detailed instructions can be found at [`get_started_for_vqav2.md`](get_started/get_started_for_vqav2.md).
|
||||
|
||||
| initialized checkpoint | resolution | augmented data | test-dev | test-std | #params | weight |
|
||||
|:----------------------------------------|:----------:|:-----:|:-----:|:-----:|:-------:|-------------------|
|
||||
| [beit3_base_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_224.pth) | 480x480 | - | 77.65 | - | 228M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_480_vqa.pth) |
|
||||
| [beit3_base_indomain_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_224.pth) | 480x480 | - | 78.46 | - | 228M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_480_vqa.pth) |
|
||||
| [beit3_large_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_224.pth) | 480x480 | - | 81.85 | - | 683M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_480_vqa.pth) |
|
||||
| [beit3_large_indomain_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_indomain_patch16_224.pth) | 480x480 | - | 82.53 | - | 683M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_indomain_patch16_480_vqa.pth) |
|
||||
| [beit3_large_indomain_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_indomain_patch16_224.pth) | 768x768 | VGQA | 82.97 | 83.03 | 684M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_indomain_patch16_768_vgqaaug_vqa.pth) |
|
||||
|
||||
|
||||
## Fine-tuning on NLVR2 (Visual Reasoning)
|
||||
|
||||
The detailed instructions can be found at [`get_started_for_nlvr2.md`](get_started/get_started_for_nlvr2.md).
|
||||
|
||||
| initialized checkpoint | resolution | dev | test-P | #params | weight |
|
||||
|:----------------------------------------|:----------:|:-----:|:-----:|:-------:|-------------------|
|
||||
| [beit3_base_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_224.pth) | 224x224 | 83.6 | 84.4 | 226M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_224_nlvr2.pth) |
|
||||
| [beit3_base_indomain_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_224.pth) | 224x224 | 84.6 | 85.3 | 226M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_224_nlvr2.pth) |
|
||||
| [beit3_large_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_224.pth) | 224x224 | 88.5 | 89.4 | 681M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_224_nlvr2.pth) |
|
||||
| [beit3_large_indomain_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_indomain_patch16_224.pth) | 224x224 | 89.2 | 90.0 | 681M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_indomain_patch16_224_nlvr2.pth) |
|
||||
|
||||
|
||||
## Fine-tuning on COCO Captioning and NoCaps (Image Captioning)
|
||||
|
||||
The detailed instructions can be found at [`get_started_for_image_captioning.md`](get_started/get_started_for_captioning.md).
|
||||
|
||||
### COCO Captioning
|
||||
|
||||
| initialized checkpoint | resolution | test CIDEr | #params | weight |
|
||||
|:----------------------------------------|:----------:|:-----:|:-------:|-------------------|
|
||||
| [beit3_base_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_224.pth) | 480x480 | 133.6 | 271M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_480_coco_captioning.pth) |
|
||||
| [beit3_base_indomain_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_224.pth) | 480x480 | 135.0 | 271M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_480_coco_captioning.pth) |
|
||||
| [beit3_large_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_224.pth) | 480x480 | 143.2 | 739M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_480_coco_captioning.pth) |
|
||||
|
||||
### NoCaps
|
||||
|
||||
| initialized checkpoint | resolution | val CIDEr | #params | weight |
|
||||
|:----------------------------------------|:----------:|:-----:|:-------:|-------------------|
|
||||
| [beit3_base_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_224.pth) | 480x480 | 104.4 | 271M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_480_nocaps.pth) |
|
||||
| [beit3_base_indomain_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_224.pth) | 480x480 | 105.6 | 271M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_indomain_patch16_480_nocaps.pth) |
|
||||
| [beit3_large_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_224.pth) | 480x480 | 120.2 | 739M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_480_nocaps.pth) |
|
||||
|
||||
|
||||
## Fine-tuning on COCO and Flickr30k Retrieval (Image-Text Retrieval)
|
||||
|
||||
The detailed instructions can be found at [`get_started_for_retrieval.md`](get_started/get_started_for_retrieval.md).
|
||||
|
||||
### COCO Retrieval
|
||||
|
||||
| initialized checkpoint | resolution | IR@1 | TR@1 | #params | weight |
|
||||
|:----------------------------------------|:----------:|:-----:|:-----:|:-------:|-------------------|
|
||||
| [beit3_base_itc_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_itc_patch16_224.pth) | 384x384 | 61.4 | 79.1 | 222M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_384_coco_retrieval.pth) |
|
||||
| [beit3_large_itc_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_itc_patch16_224.pth) | 384x384 | 63.4 | 82.1 | 675M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_384_coco_retrieval.pth) |
|
||||
|
||||
### Flickr30k Retrieval
|
||||
|
||||
| initialized checkpoint | resolution | IR@1 | TR@1 | #params | weight |
|
||||
|:----------------------------------------|:----------:|:-----:|:-----:|:-------:|-------------------|
|
||||
| [beit3_base_itc_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_base_itc_patch16_224.pth) | 384x384 | 86.2 | 96.3 | 222M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_base_patch16_384_f30k_retrieval.pth) |
|
||||
| [beit3_large_itc_patch16_224](https://github.com/addf400/files/releases/download/beit3/beit3_large_itc_patch16_224.pth) | 384x384 | 88.1 | 97.2 | 675M | [link](https://github.com/addf400/files/releases/download/beit3/beit3_large_patch16_384_f30k_retrieval.pth) |
|
||||
|
||||
|
||||
## Citation
|
||||
|
||||
If you find this repository useful, please consider citing our work:
|
||||
```
|
||||
@inproceedings{beit3,
|
||||
title={Image as a foreign language: {BEiT} pretraining for vision and vision-language tasks},
|
||||
author={Wenhui Wang and Hangbo Bao and Li Dong and Johan Bjorck and Zhiliang Peng and Qiang Liu and Kriti Aggarwal and Owais Khan Mohammed and Saksham Singhal and Subhojit Som and Furu Wei},
|
||||
booktitle={Proceedings of the IEEE/CVF Conference on Computer Vision and Pattern Recognition},
|
||||
year={2023}
|
||||
}
|
||||
|
||||
@article{beitv2,
|
||||
title={{BEiT v2}: Masked Image Modeling with Vector-Quantized Visual Tokenizers},
|
||||
author={Zhiliang Peng and Li Dong and Hangbo Bao and Qixiang Ye and Furu Wei},
|
||||
year={2022},
|
||||
eprint={2208.06366},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.CV}
|
||||
}
|
||||
|
||||
@inproceedings{beit,
|
||||
title={{BEiT}: {BERT} Pre-Training of Image Transformers},
|
||||
author={Hangbo Bao and Li Dong and Songhao Piao and Furu Wei},
|
||||
booktitle={International Conference on Learning Representations},
|
||||
year={2022},
|
||||
url={https://openreview.net/forum?id=p-BhZSz59o4}
|
||||
}
|
||||
```
|
||||
|
||||
|
||||
## Acknowledgement
|
||||
|
||||
This repository is built using the [BEiT](https://github.com/microsoft/unilm/tree/master/beit), the [BEiTv2](https://github.com/microsoft/unilm/tree/master/beit2), the [CLIP](https://github.com/openai/CLIP), the [open_clip](https://github.com/mlfoundations/open_clip), the [Oscar](https://github.com/microsoft/Oscar), the [DeiT](https://github.com/facebookresearch/deit), the [Dino](https://github.com/facebookresearch/dino) repository and the [timm](https://github.com/rwightman/pytorch-image-models) library.
|
||||
|
||||
|
||||
## License
|
||||
This project is licensed under the license found in the LICENSE file in the root directory of this source tree.
|
||||
|
||||
[Microsoft Open Source Code of Conduct](https://opensource.microsoft.com/codeofconduct)
|
||||
|
||||
### Contact Information
|
||||
|
||||
For help or issues using BEiT-3 models, please submit a GitHub issue.
|
||||
@@ -0,0 +1,847 @@
|
||||
# --------------------------------------------------------
|
||||
# Image as a Foreign Language: BEiT Pretraining for Vision and Vision-Language Tasks (https://arxiv.org/abs/2208.10442)
|
||||
# Github source: https://github.com/microsoft/unilm/tree/master/beit3
|
||||
# Copyright (c) 2023 Microsoft
|
||||
# Licensed under The MIT License [see LICENSE for details]
|
||||
# --------------------------------------------------------'
|
||||
|
||||
import os
|
||||
import json
|
||||
import random
|
||||
import torch
|
||||
import glob
|
||||
from collections import defaultdict, Counter
|
||||
from torchvision import transforms
|
||||
from torchvision.datasets.folder import default_loader
|
||||
from timm.data.constants import IMAGENET_DEFAULT_MEAN, IMAGENET_DEFAULT_STD, IMAGENET_INCEPTION_MEAN, IMAGENET_INCEPTION_STD
|
||||
from timm.data.transforms import RandomResizedCropAndInterpolation
|
||||
from timm.data import create_transform
|
||||
|
||||
import utils
|
||||
from glossary import normalize_word
|
||||
from randaug import RandomAugment
|
||||
|
||||
|
||||
class BaseDataset(torch.utils.data.Dataset):
|
||||
def __init__(
|
||||
self, data_path, split, transform,
|
||||
tokenizer, num_max_bpe_tokens, task=None,
|
||||
):
|
||||
index_files = self.get_index_files(split, task=task)
|
||||
self.tokenizer = tokenizer
|
||||
self.num_max_bpe_tokens = num_max_bpe_tokens
|
||||
self.data_path = data_path
|
||||
items = []
|
||||
self.index_files = index_files
|
||||
|
||||
offset = 0
|
||||
for _index_file in index_files:
|
||||
index_file = os.path.join(data_path, _index_file)
|
||||
with open(index_file, mode="r", encoding="utf-8") as reader:
|
||||
for line in reader:
|
||||
data = json.loads(line)
|
||||
items.append(data)
|
||||
print("Load %d image-text pairs from %s. " % (len(items) - offset, index_file))
|
||||
offset = len(items)
|
||||
self.items = items
|
||||
self.bos_token_id = tokenizer.bos_token_id
|
||||
self.eos_token_id = tokenizer.eos_token_id
|
||||
self.pad_token_id = tokenizer.pad_token_id
|
||||
self.loader = default_loader
|
||||
self.transform = transform
|
||||
self.split = split
|
||||
|
||||
@staticmethod
|
||||
def get_index_files(split):
|
||||
raise NotImplementedError()
|
||||
|
||||
def _get_image(self, image_path: str):
|
||||
image_path = os.path.join(self.data_path, image_path)
|
||||
image = self.loader(image_path)
|
||||
return self.transform(image)
|
||||
|
||||
def _get_text_segment(self, text_segment, max_len=None):
|
||||
if isinstance(text_segment, str):
|
||||
tokens = self.tokenizer.tokenize(text_segment)
|
||||
else:
|
||||
tokens = text_segment[:]
|
||||
if len(tokens) == 0:
|
||||
raise RuntimeError("The text segment should contains at least one tokens!")
|
||||
if max_len is None:
|
||||
max_len = self.num_max_bpe_tokens
|
||||
|
||||
if len(tokens) > max_len - 2:
|
||||
tokens = tokens[:max_len - 2]
|
||||
|
||||
tokens = [self.bos_token_id] + tokens[:] + [self.eos_token_id]
|
||||
num_tokens = len(tokens)
|
||||
padding_mask = [0] * num_tokens + [1] * (max_len - num_tokens)
|
||||
return tokens + [self.pad_token_id] * (max_len - num_tokens), padding_mask, num_tokens
|
||||
|
||||
def _get_image_text_example(self, index: int, data: dict):
|
||||
item = self.items[index]
|
||||
img_path = item["image_path"]
|
||||
img = self._get_image(img_path)
|
||||
data["image"] = img
|
||||
|
||||
text_segment = item["text_segment"]
|
||||
language_tokens, padding_mask, _ = self._get_text_segment(text_segment)
|
||||
data["language_tokens"] = language_tokens
|
||||
data["padding_mask"] = padding_mask
|
||||
|
||||
def __getitem__(self, index: int):
|
||||
data = dict()
|
||||
self._get_image_text_example(index, data)
|
||||
return data
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.items)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
head = "Dataset " + self.__class__.__name__
|
||||
body = '{' + "\n Number of items: %s," % self.__len__()
|
||||
body += "\n data root = %s," % self.data_path
|
||||
body += "\n split = %s," % self.split
|
||||
body += "\n dataset index files = %s" % str(self.index_files)
|
||||
body += "\n num max bpe tokens = %s" % self.num_max_bpe_tokens
|
||||
body += "\n transforms = ["
|
||||
for t in self.transform.transforms:
|
||||
body += "\n %s" % str(t)
|
||||
body += "\n ]"
|
||||
body += "\n}"
|
||||
|
||||
return head + body
|
||||
|
||||
|
||||
def _write_data_into_jsonl(items, jsonl_file):
|
||||
with open(jsonl_file, mode="w", encoding="utf-8") as writer:
|
||||
for data in items:
|
||||
writer.write(json.dumps(data, indent=None))
|
||||
writer.write('\n')
|
||||
print("Write %s with %d items !" % (jsonl_file, len(items)))
|
||||
|
||||
|
||||
def _make_retrieval_coco_karpathy_dataset_index(
|
||||
data_path,
|
||||
tokenizer,
|
||||
split=("train", "restval"),
|
||||
split_name="train",
|
||||
):
|
||||
coco_karpathy_split_json_file = os.path.join(data_path, "dataset_coco.json")
|
||||
items = []
|
||||
image_counter = set()
|
||||
print("read %s" % coco_karpathy_split_json_file)
|
||||
with open(coco_karpathy_split_json_file, mode="r", encoding="utf-8") as reader:
|
||||
data = json.loads(reader.read())
|
||||
for item in data["images"]:
|
||||
if item["split"] in split:
|
||||
image_path = os.path.join(item["filepath"], item["filename"])
|
||||
for sent in item["sentences"]:
|
||||
tokens = tokenizer.tokenize(sent["raw"])
|
||||
token_ids = tokenizer.convert_tokens_to_ids(tokens)
|
||||
items.append({
|
||||
"image_path": image_path,
|
||||
"text_segment": token_ids,
|
||||
"image_id": len(image_counter),
|
||||
})
|
||||
if image_path not in image_counter:
|
||||
image_counter.add(image_path)
|
||||
print("Find %d images and %d image-text pairs for karpathy dataset %s split !" % \
|
||||
(len(image_counter), len(items), split_name))
|
||||
index_file = os.path.join(data_path, "coco_retrieval.%s.jsonl" % split_name)
|
||||
_write_data_into_jsonl(items, index_file)
|
||||
pass
|
||||
|
||||
|
||||
def _make_captioning_coco_karpathy_dataset_index(
|
||||
data_path,
|
||||
tokenizer,
|
||||
split=("train", "restval"),
|
||||
split_name="train",
|
||||
):
|
||||
coco_karpathy_split_json_file = os.path.join(data_path, "dataset_coco.json")
|
||||
items = []
|
||||
image_counter = set()
|
||||
print("read %s" % coco_karpathy_split_json_file)
|
||||
with open(coco_karpathy_split_json_file, mode="r", encoding="utf-8") as reader:
|
||||
data = json.loads(reader.read())
|
||||
for item in data["images"]:
|
||||
if item["split"] in split:
|
||||
image_path = os.path.join(item["filepath"], item["filename"])
|
||||
if item["split"] in ["train", "restval"]:
|
||||
for sent in item["sentences"]:
|
||||
tokens = tokenizer.tokenize(sent["raw"])
|
||||
token_ids = tokenizer.convert_tokens_to_ids(tokens)
|
||||
items.append({
|
||||
"image_path": image_path,
|
||||
"text_segment": token_ids,
|
||||
"image_id": item["cocoid"],
|
||||
})
|
||||
else:
|
||||
items.append({
|
||||
"image_path": image_path,
|
||||
"text_segment": None,
|
||||
"image_id": item["cocoid"],
|
||||
})
|
||||
if image_path not in image_counter:
|
||||
image_counter.add(image_path)
|
||||
print("Find %d images and %d image-text pairs for karpathy dataset %s split !" % \
|
||||
(len(image_counter), len(items), split_name))
|
||||
index_file = os.path.join(data_path, "coco_captioning.%s.jsonl" % split_name)
|
||||
_write_data_into_jsonl(items, index_file)
|
||||
pass
|
||||
|
||||
|
||||
def _make_nocaps_dataset_index(
|
||||
data_path,
|
||||
split="val",
|
||||
):
|
||||
if split == "val":
|
||||
json_file = "nocaps_val_4500_captions.json"
|
||||
elif split == "test":
|
||||
json_file = "nocaps_test_image_info.json"
|
||||
nocaps_split_json_file = os.path.join(data_path, json_file)
|
||||
items = []
|
||||
image_counter = set()
|
||||
print("read %s" % nocaps_split_json_file)
|
||||
with open(nocaps_split_json_file, mode="r", encoding="utf-8") as reader:
|
||||
data = json.loads(reader.read())
|
||||
for item in data["images"]:
|
||||
image_path = os.path.join(split, item["file_name"])
|
||||
items.append({
|
||||
"image_path": image_path,
|
||||
"text_segment": None,
|
||||
"image_id": item["id"],
|
||||
})
|
||||
|
||||
if image_path not in image_counter:
|
||||
image_counter.add(image_path)
|
||||
|
||||
print("Find %d images and %d image-text pairs for nocaps dataset %s split !" % \
|
||||
(len(image_counter), len(items), split))
|
||||
index_file = os.path.join(data_path, "nocaps.%s.jsonl" % split)
|
||||
_write_data_into_jsonl(items, index_file)
|
||||
|
||||
|
||||
class NLVR2Dataset(BaseDataset):
|
||||
@staticmethod
|
||||
def get_index_files(split, task=None):
|
||||
if split == "train":
|
||||
return ("nlvr2.train.index.jsonl", )
|
||||
elif split == "val":
|
||||
return ("nlvr2.dev.index.jsonl", )
|
||||
elif split == "test":
|
||||
return ("nlvr2.test-P.index.jsonl", )
|
||||
else:
|
||||
raise RuntimeError("split %s is not found!" % split)
|
||||
|
||||
def __getitem__(self, index: int):
|
||||
data = super().__getitem__(index)
|
||||
item = self.items[index]
|
||||
img_path = item["image2_path"]
|
||||
img = self._get_image(img_path)
|
||||
data["image2"] = img
|
||||
data["label"] = self.items[index]["label"]
|
||||
return data
|
||||
|
||||
@staticmethod
|
||||
def __preprocess_json(preifx, json_file, tokenizer, index_file):
|
||||
items = []
|
||||
with open(json_file, mode="r", encoding="utf-8") as reader:
|
||||
for line in reader:
|
||||
data = json.loads(line)
|
||||
path = os.path.join(preifx, str(data["directory"])) if "directory" in data else preifx
|
||||
path = os.path.join(path, "-".join(data["identifier"].split("-")[:-1]))
|
||||
tokens = tokenizer.tokenize(data["sentence"])
|
||||
token_ids = tokenizer.convert_tokens_to_ids(tokens)
|
||||
items.append({
|
||||
"image_path": path + "-img0.png",
|
||||
"image2_path": path + "-img1.png",
|
||||
"text_segment": token_ids,
|
||||
"label": 1 if data["label"] == "True" else 0,
|
||||
"identifier": data["identifier"],
|
||||
})
|
||||
_write_data_into_jsonl(items, index_file)
|
||||
|
||||
@classmethod
|
||||
def make_dataset_index(cls, data_path, tokenizer, nlvr_repo_path):
|
||||
cls.__preprocess_json(
|
||||
preifx="images/train", json_file=os.path.join(nlvr_repo_path, "nlvr2/data/train.json"),
|
||||
tokenizer=tokenizer, index_file=os.path.join(data_path, cls.get_index_files("train")[0]),
|
||||
)
|
||||
cls.__preprocess_json(
|
||||
preifx="dev", json_file=os.path.join(nlvr_repo_path, "nlvr2/data/dev.json"),
|
||||
tokenizer=tokenizer, index_file=os.path.join(data_path, cls.get_index_files("val")[0]),
|
||||
)
|
||||
cls.__preprocess_json(
|
||||
preifx="test1", json_file=os.path.join(nlvr_repo_path, "nlvr2/data/test1.json"),
|
||||
tokenizer=tokenizer, index_file=os.path.join(data_path, cls.get_index_files("test")[0]),
|
||||
)
|
||||
|
||||
|
||||
class ImageNetDataset(BaseDataset):
|
||||
@staticmethod
|
||||
def get_index_files(split, task=None):
|
||||
if split == "train":
|
||||
return ("imagenet.train.index.jsonl", )
|
||||
elif split == "val":
|
||||
return ("imagenet.val.index.jsonl", )
|
||||
elif split == "test":
|
||||
return ("imagenet.val.index.jsonl", )
|
||||
else:
|
||||
raise RuntimeError("split %s is not found!" % split)
|
||||
|
||||
def __getitem__(self, index: int):
|
||||
data = dict()
|
||||
item = self.items[index]
|
||||
img_path = item["image_path"]
|
||||
img = self._get_image(img_path)
|
||||
data["image"] = img
|
||||
data["label"] = item["label"]
|
||||
return data
|
||||
|
||||
@staticmethod
|
||||
def _find_classes(dir):
|
||||
"""
|
||||
Finds the class folders in a dataset.
|
||||
Args:
|
||||
dir (string): Root directory path.
|
||||
Returns:
|
||||
tuple: (classes, class_to_idx) where classes are relative to (dir), and class_to_idx is a dictionary.
|
||||
Ensures:
|
||||
No class is a subdirectory of another.
|
||||
"""
|
||||
classes = [d.name for d in os.scandir(dir) if d.is_dir()]
|
||||
classes.sort()
|
||||
class_to_idx = {cls_name: i for i, cls_name in enumerate(classes)}
|
||||
return classes, class_to_idx
|
||||
|
||||
@staticmethod
|
||||
def _make_imagenet_index(data_path, index_path, data_path_prefix, class_to_idx, split):
|
||||
items = []
|
||||
index_file = os.path.join(index_path, f"imagenet.{split}.index.jsonl")
|
||||
for target_class in sorted(class_to_idx.keys()):
|
||||
class_index = class_to_idx[target_class]
|
||||
target_dir = os.path.join(data_path, target_class)
|
||||
if not os.path.isdir(target_dir):
|
||||
continue
|
||||
for root, _, fnames in sorted(os.walk(target_dir, followlinks=True)):
|
||||
for fname in sorted(fnames):
|
||||
path = os.path.join(root, fname)
|
||||
path = path.replace(data_path_prefix, "")
|
||||
items.append({
|
||||
"image_path": path,
|
||||
"label": class_index,
|
||||
})
|
||||
|
||||
_write_data_into_jsonl(items, index_file)
|
||||
|
||||
@classmethod
|
||||
def make_dataset_index(cls, train_data_path, val_data_path, index_path):
|
||||
data_path_prefix = train_data_path[:[x[0]==x[1] for x in zip(train_data_path, val_data_path)].index(0)]
|
||||
classes, class_to_idx = cls._find_classes(train_data_path)
|
||||
cls._make_imagenet_index(
|
||||
data_path=train_data_path, index_path=index_path, data_path_prefix=data_path_prefix,
|
||||
class_to_idx=class_to_idx, split="train",
|
||||
)
|
||||
cls._make_imagenet_index(
|
||||
data_path=val_data_path, index_path=index_path, data_path_prefix=data_path_prefix,
|
||||
class_to_idx=class_to_idx, split="val",
|
||||
)
|
||||
|
||||
|
||||
class VQAv2Dataset(BaseDataset):
|
||||
def __init__(self, data_path, **kwargs):
|
||||
super().__init__(data_path=data_path, **kwargs)
|
||||
ans2label_file = os.path.join(data_path, "answer2label.txt")
|
||||
ans2label = {}
|
||||
label2ans = []
|
||||
with open(ans2label_file, mode="r", encoding="utf-8") as reader:
|
||||
for i, line in enumerate(reader):
|
||||
data = json.loads(line)
|
||||
ans = data["answer"]
|
||||
label = data["label"]
|
||||
label = int(label)
|
||||
assert label == i
|
||||
ans2label[ans] = i
|
||||
label2ans.append(ans)
|
||||
|
||||
self.ans2label = ans2label
|
||||
self.label2ans = label2ans
|
||||
|
||||
@staticmethod
|
||||
def get_index_files(split, task=None):
|
||||
if split == "train":
|
||||
return ("vqa.train.jsonl", "vqa.trainable_val.jsonl")
|
||||
elif split == "val":
|
||||
return ("vqa.rest_val.jsonl", )
|
||||
elif split == "test":
|
||||
return ("vqa.test.jsonl", )
|
||||
elif split == "test-dev":
|
||||
return ("vqa.test-dev.jsonl", )
|
||||
else:
|
||||
raise RuntimeError("split %s is not found!" % split)
|
||||
|
||||
def __getitem__(self, index: int):
|
||||
data = super().__getitem__(index)
|
||||
if "labels" in self.items[index] and len(self.items[index]["labels"]) > 0:
|
||||
labels = [0.] * len(self.label2ans)
|
||||
for l, s in zip(self.items[index]["labels"], self.items[index]["scores"]):
|
||||
labels[l] = s
|
||||
data["labels"] = torch.FloatTensor(labels)
|
||||
else:
|
||||
data["qid"] = self.items[index]["qid"]
|
||||
return data
|
||||
|
||||
@staticmethod
|
||||
def get_score(occurences):
|
||||
if occurences == 0:
|
||||
return 0.0
|
||||
elif occurences == 1:
|
||||
return 0.3
|
||||
elif occurences == 2:
|
||||
return 0.6
|
||||
elif occurences == 3:
|
||||
return 0.9
|
||||
else:
|
||||
return 1.0
|
||||
|
||||
@classmethod
|
||||
def make_dataset_index(cls, data_path, tokenizer, annotation_data_path):
|
||||
with open(os.path.join(annotation_data_path, "v2_OpenEnded_mscoco_train2014_questions.json"), "r") as fp:
|
||||
questions_train2014 = json.load(fp)["questions"]
|
||||
with open(os.path.join(annotation_data_path, "v2_OpenEnded_mscoco_val2014_questions.json"), "r") as fp:
|
||||
questions_val2014 = json.load(fp)["questions"]
|
||||
with open(os.path.join(annotation_data_path, "v2_OpenEnded_mscoco_test2015_questions.json"), "r") as fp:
|
||||
questions_test2015 = json.load(fp)["questions"]
|
||||
with open(os.path.join(annotation_data_path, "v2_OpenEnded_mscoco_test-dev2015_questions.json"), "r") as fp:
|
||||
questions_test_dev2015 = json.load(fp)["questions"]
|
||||
|
||||
with open(os.path.join(annotation_data_path, "v2_mscoco_train2014_annotations.json"), "r") as fp:
|
||||
annotations_train2014 = json.load(fp)["annotations"]
|
||||
with open(os.path.join(annotation_data_path, "v2_mscoco_val2014_annotations.json"), "r") as fp:
|
||||
annotations_val2014 = json.load(fp)["annotations"]
|
||||
|
||||
annotations = dict()
|
||||
|
||||
for split, questions in zip(
|
||||
["train", "val", "test", "test-dev"],
|
||||
[questions_train2014, questions_val2014, questions_test2015, questions_test_dev2015],
|
||||
):
|
||||
_annot = defaultdict(dict)
|
||||
for q in questions:
|
||||
question_text = q["question"]
|
||||
tokens = tokenizer.tokenize(question_text)
|
||||
token_ids = tokenizer.convert_tokens_to_ids(tokens)
|
||||
|
||||
assert q["question_id"] not in _annot[q["image_id"]]
|
||||
_annot[q["image_id"]][q["question_id"]] = {
|
||||
"question": question_text,
|
||||
"token_ids": token_ids,
|
||||
}
|
||||
|
||||
annotations[split] = _annot
|
||||
|
||||
all_major_answers = list()
|
||||
|
||||
for split, annots in zip(
|
||||
["train", "val"], [annotations_train2014, annotations_val2014],
|
||||
):
|
||||
# _annot = annotations[split]
|
||||
for q in annots:
|
||||
all_major_answers.append(q["multiple_choice_answer"])
|
||||
|
||||
all_major_answers = [normalize_word(word) for word in all_major_answers]
|
||||
counter = {k: v for k, v in Counter(all_major_answers).items() if v >= 9}
|
||||
ans2label = {k: i for i, k in enumerate(counter.keys())}
|
||||
label2ans = list(counter.keys())
|
||||
|
||||
for split, annots in zip(
|
||||
["train", "val"], [annotations_train2014, annotations_val2014],
|
||||
):
|
||||
_annot = annotations[split]
|
||||
for q in annots:
|
||||
answers = q["answers"]
|
||||
answer_count = {}
|
||||
for answer in answers:
|
||||
answer_ = answer["answer"]
|
||||
answer_count[answer_] = answer_count.get(answer_, 0) + 1
|
||||
|
||||
labels = []
|
||||
scores = []
|
||||
for answer in answer_count:
|
||||
if answer not in ans2label:
|
||||
continue
|
||||
labels.append(ans2label[answer])
|
||||
score = cls.get_score(answer_count[answer])
|
||||
scores.append(score)
|
||||
|
||||
assert "labels" not in _annot[q["image_id"]][q["question_id"]]
|
||||
assert "question" in _annot[q["image_id"]][q["question_id"]]
|
||||
_annot[q["image_id"]][q["question_id"]]["labels"] = labels
|
||||
_annot[q["image_id"]][q["question_id"]]["scores"] = scores
|
||||
|
||||
for split in ["train", "val"]:
|
||||
filtered_annot = dict()
|
||||
for ik, iv in annotations[split].items():
|
||||
new_q = dict()
|
||||
for qk, qv in iv.items():
|
||||
if len(qv["labels"]) != 0:
|
||||
new_q[qk] = qv
|
||||
if len(new_q) != 0:
|
||||
filtered_annot[ik] = new_q
|
||||
annotations[split] = filtered_annot
|
||||
|
||||
split2items = {}
|
||||
for split in ["train", "val", "test", "test-dev"]:
|
||||
annot = annotations[split]
|
||||
split_name = {
|
||||
"train": "train2014",
|
||||
"val": "val2014",
|
||||
"test": "test2015",
|
||||
"test-dev": "test2015",
|
||||
}[split]
|
||||
paths = list(glob.glob(f"{data_path}/{split_name}/*.jpg"))
|
||||
random.shuffle(paths)
|
||||
annot_paths = [path for path in paths \
|
||||
if int(path.split("/")[-1].split("_")[-1][:-4]) in annot]
|
||||
|
||||
if len(paths) == len(annot_paths):
|
||||
print("all images have caption annotations")
|
||||
else:
|
||||
print("not all images have caption annotations")
|
||||
print(len(paths), len(annot_paths), len(annot))
|
||||
|
||||
items = []
|
||||
for path in annot_paths:
|
||||
iid = int(path.split("/")[-1].split("_")[-1][:-4])
|
||||
_annot = annotations[split][iid]
|
||||
for qid in _annot:
|
||||
q = _annot[qid]
|
||||
if split in ["train", "val"]:
|
||||
labels = q["labels"]
|
||||
scores = q["scores"]
|
||||
else:
|
||||
labels, scores = [], []
|
||||
|
||||
items.append({
|
||||
"image_path": os.path.join(split_name, path.split('/')[-1]),
|
||||
"text_segment": q["token_ids"],
|
||||
"labels": labels,
|
||||
"scores": scores,
|
||||
"qid": qid,
|
||||
})
|
||||
split2items[split] = items
|
||||
|
||||
_write_data_into_jsonl(items=items, jsonl_file=os.path.join(data_path, "vqa.%s.jsonl" % split))
|
||||
|
||||
# Following ViLT, we use 1000 images of the original val set as the final val set
|
||||
val_image2items = defaultdict(list)
|
||||
for item in split2items["val"]:
|
||||
val_image2items[item["image_path"]].append(item)
|
||||
|
||||
print("Contains %d image and %d pairs for val set!" % (len(val_image2items), len(split2items["val"])))
|
||||
|
||||
val_images = list(val_image2items.keys())
|
||||
random.shuffle(val_images)
|
||||
trainable_val = []
|
||||
rest_val = []
|
||||
for i, image_id in enumerate(val_images):
|
||||
if i < 1000:
|
||||
rest_val += val_image2items[image_id]
|
||||
else:
|
||||
trainable_val += val_image2items[image_id]
|
||||
|
||||
_write_data_into_jsonl(items=trainable_val, jsonl_file=os.path.join(data_path, "vqa.trainable_val.jsonl"))
|
||||
_write_data_into_jsonl(items=rest_val, jsonl_file=os.path.join(data_path, "vqa.rest_val.jsonl"))
|
||||
|
||||
with open(os.path.join(data_path, "answer2label.txt"), mode="w", encoding="utf-8") as writer:
|
||||
for ans in ans2label:
|
||||
to_json = {
|
||||
"answer": ans,
|
||||
"label": ans2label[ans]
|
||||
}
|
||||
writer.write("%s\n" % json.dumps(to_json))
|
||||
|
||||
|
||||
class RetrievalDataset(BaseDataset):
|
||||
@staticmethod
|
||||
def get_index_files(split, task=None):
|
||||
if split == "train":
|
||||
return (f"{task}.train.jsonl", )
|
||||
elif split == "val":
|
||||
return (f"{task}.val.jsonl", )
|
||||
elif split == "test":
|
||||
return (f"{task}.test.jsonl", )
|
||||
else:
|
||||
raise RuntimeError("split %s is not found!" % split)
|
||||
|
||||
def __getitem__(self, index: int):
|
||||
data = super().__getitem__(index)
|
||||
data["image_id"] = self.items[index]["image_id"]
|
||||
return data
|
||||
|
||||
@staticmethod
|
||||
def make_flickr30k_dataset_index(data_path, tokenizer, karpathy_path):
|
||||
|
||||
with open(os.path.join(karpathy_path, "dataset_flickr30k.json"), "r") as reader:
|
||||
captions = json.loads(reader.read())
|
||||
|
||||
captions = captions["images"]
|
||||
split2items = defaultdict(list)
|
||||
split2images = defaultdict(set)
|
||||
|
||||
for each_item in captions:
|
||||
image_path = os.path.join("flickr30k-images", each_item["filename"])
|
||||
split = each_item["split"]
|
||||
|
||||
for text_segment in each_item["sentences"]:
|
||||
tokens = tokenizer.tokenize(text_segment["raw"])
|
||||
token_ids = tokenizer.convert_tokens_to_ids(tokens)
|
||||
|
||||
split2items[split].append({
|
||||
"image_path": image_path,
|
||||
"text_segment": token_ids,
|
||||
"image_id": len(split2images[split]),
|
||||
})
|
||||
|
||||
assert each_item["filename"] not in split2images[split]
|
||||
split2images[split].add(each_item["filename"])
|
||||
|
||||
for split in split2items:
|
||||
print("%d images and %d image-text pairs!" % (len(split2images[split]), len(split2items[split])))
|
||||
_write_data_into_jsonl(split2items[split], os.path.join(data_path, "flickr30k.%s.jsonl" % split))
|
||||
|
||||
@staticmethod
|
||||
def make_coco_dataset_index(data_path, tokenizer):
|
||||
_make_retrieval_coco_karpathy_dataset_index(data_path, tokenizer, split=("train", "restval"), split_name="train")
|
||||
_make_retrieval_coco_karpathy_dataset_index(data_path, tokenizer, split=("val", ), split_name="val")
|
||||
_make_retrieval_coco_karpathy_dataset_index(data_path, tokenizer, split=("test", ), split_name="test")
|
||||
|
||||
|
||||
class CaptioningDataset(BaseDataset):
|
||||
|
||||
def __init__(self, data_path, split, transform,
|
||||
tokenizer, num_max_bpe_tokens, task, mask_prob):
|
||||
super().__init__(
|
||||
data_path=data_path, split=split,
|
||||
transform=transform, tokenizer=tokenizer,
|
||||
num_max_bpe_tokens=num_max_bpe_tokens, task=task,
|
||||
)
|
||||
self.mask_token_id = tokenizer.mask_token_id
|
||||
self.language_vocab_size = tokenizer.vocab_size
|
||||
self.mask_prob = mask_prob
|
||||
|
||||
@staticmethod
|
||||
def get_index_files(split, task=None):
|
||||
if split == "train":
|
||||
return ("coco_captioning.train.jsonl", )
|
||||
elif split == "val":
|
||||
return (f"{task}.val.jsonl", )
|
||||
elif split == "test":
|
||||
return (f"{task}.test.jsonl", )
|
||||
else:
|
||||
raise RuntimeError("split %s is not found!" % split)
|
||||
|
||||
def _get_mask_token(self, token):
|
||||
p = random.random()
|
||||
if p < 0.8:
|
||||
return self.mask_token_id
|
||||
elif p < 0.9:
|
||||
return token
|
||||
else:
|
||||
return random.randint(3, self.language_vocab_size - 1)
|
||||
|
||||
def _masking_on_text_tokens(self, tokens, num_tokens, mask_prob):
|
||||
bool_masked_pos = [0] * len(tokens)
|
||||
to_mask = min(int(num_tokens * mask_prob + 0.5), num_tokens - 1)
|
||||
to_mask = max(to_mask, 1)
|
||||
num_masked_tokens = 0
|
||||
while num_masked_tokens < to_mask:
|
||||
i = random.randint(1, num_tokens - 1)
|
||||
if bool_masked_pos[i] == 0:
|
||||
bool_masked_pos[i] = 1
|
||||
tokens[i] = self._get_mask_token(tokens[i])
|
||||
num_masked_tokens += 1
|
||||
|
||||
return tokens, bool_masked_pos
|
||||
|
||||
def __getitem__(self, index: int):
|
||||
data = dict()
|
||||
item = self.items[index]
|
||||
img_path = item["image_path"]
|
||||
img = self._get_image(img_path)
|
||||
data["image"] = img
|
||||
data["image_id"] = item["image_id"]
|
||||
|
||||
text_segment = item["text_segment"]
|
||||
if text_segment is not None:
|
||||
language_tokens, padding_mask, num_tokens = self._get_text_segment(text_segment)
|
||||
masked_tokens = language_tokens[:]
|
||||
masked_tokens, language_masked_pos = \
|
||||
self._masking_on_text_tokens(masked_tokens, num_tokens, self.mask_prob)
|
||||
data["language_tokens"] = language_tokens
|
||||
data["masked_tokens"] = masked_tokens
|
||||
data["language_masked_pos"] = language_masked_pos
|
||||
data["padding_mask"] = padding_mask
|
||||
return data
|
||||
|
||||
@staticmethod
|
||||
def make_coco_captioning_dataset_index(data_path, tokenizer):
|
||||
_make_captioning_coco_karpathy_dataset_index(data_path, tokenizer, split=("train", "restval"), split_name="train")
|
||||
_make_captioning_coco_karpathy_dataset_index(data_path, tokenizer, split=("val", ), split_name="val")
|
||||
_make_captioning_coco_karpathy_dataset_index(data_path, tokenizer, split=("test", ), split_name="test")
|
||||
|
||||
@staticmethod
|
||||
def make_nocaps_captioning_dataset_index(data_path):
|
||||
_make_nocaps_dataset_index(data_path, split="val")
|
||||
_make_nocaps_dataset_index(data_path, split="test")
|
||||
|
||||
|
||||
task2dataset = {
|
||||
"nlvr2": NLVR2Dataset,
|
||||
"vqav2": VQAv2Dataset,
|
||||
"flickr30k": RetrievalDataset,
|
||||
"coco_retrieval": RetrievalDataset,
|
||||
"coco_captioning": CaptioningDataset,
|
||||
"nocaps": CaptioningDataset,
|
||||
"imagenet": ImageNetDataset,
|
||||
}
|
||||
|
||||
|
||||
def create_dataloader(dataset, is_train, batch_size, num_workers, pin_mem, dist_eval=False):
|
||||
if is_train or dist_eval:
|
||||
num_tasks = utils.get_world_size()
|
||||
global_rank = utils.get_rank()
|
||||
|
||||
if not is_train and dist_eval and len(dataset) % num_tasks != 0:
|
||||
print('Warning: Enabling distributed evaluation with an eval dataset not divisible by process number. '
|
||||
'This will slightly alter validation results as extra duplicate entries are added to achieve '
|
||||
'equal num of samples per-process.')
|
||||
|
||||
sampler = torch.utils.data.DistributedSampler(
|
||||
dataset, num_replicas=num_tasks, rank=global_rank, shuffle=is_train
|
||||
)
|
||||
else:
|
||||
sampler = torch.utils.data.SequentialSampler(dataset)
|
||||
|
||||
return torch.utils.data.DataLoader(
|
||||
dataset, sampler=sampler,
|
||||
batch_size=batch_size,
|
||||
num_workers=num_workers,
|
||||
pin_memory=pin_mem,
|
||||
drop_last=is_train,
|
||||
collate_fn=utils.merge_batch_tensors_by_dict_key,
|
||||
)
|
||||
|
||||
|
||||
def build_transform(is_train, args):
|
||||
if args.task in ["imagenet"]:
|
||||
return build_imagenet_transform(is_train, args)
|
||||
|
||||
if is_train:
|
||||
t = [
|
||||
RandomResizedCropAndInterpolation(args.input_size, scale=(0.5, 1.0), interpolation=args.train_interpolation),
|
||||
transforms.RandomHorizontalFlip(),
|
||||
]
|
||||
if args.randaug:
|
||||
t.append(
|
||||
RandomAugment(
|
||||
2, 7, isPIL=True,
|
||||
augs=[
|
||||
'Identity','AutoContrast','Equalize','Brightness','Sharpness',
|
||||
'ShearX', 'ShearY', 'TranslateX', 'TranslateY', 'Rotate',
|
||||
]))
|
||||
t += [
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize(mean=IMAGENET_INCEPTION_MEAN, std=IMAGENET_INCEPTION_STD),
|
||||
]
|
||||
t = transforms.Compose(t)
|
||||
else:
|
||||
t = transforms.Compose([
|
||||
transforms.Resize((args.input_size, args.input_size), interpolation=3),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize(mean=IMAGENET_INCEPTION_MEAN, std=IMAGENET_INCEPTION_STD)
|
||||
])
|
||||
|
||||
return t
|
||||
|
||||
|
||||
def build_imagenet_transform(is_train, args):
|
||||
resize_im = args.input_size > 32
|
||||
if is_train:
|
||||
# this should always dispatch to transforms_imagenet_train
|
||||
transform = create_transform(
|
||||
input_size=args.input_size,
|
||||
is_training=True,
|
||||
color_jitter=args.color_jitter,
|
||||
auto_augment=args.aa,
|
||||
interpolation=args.train_interpolation,
|
||||
re_prob=args.reprob,
|
||||
re_mode=args.remode,
|
||||
re_count=args.recount,
|
||||
mean=IMAGENET_DEFAULT_MEAN,
|
||||
std=IMAGENET_DEFAULT_STD,
|
||||
)
|
||||
if not resize_im:
|
||||
# replace RandomResizedCropAndInterpolation with
|
||||
# RandomCrop
|
||||
transform.transforms[0] = transforms.RandomCrop(
|
||||
args.input_size, padding=4)
|
||||
return transform
|
||||
|
||||
t = []
|
||||
if resize_im:
|
||||
if args.crop_pct is None:
|
||||
args.crop_pct = 1.0
|
||||
size = int(args.input_size / args.crop_pct)
|
||||
t.append(
|
||||
transforms.Resize(size, interpolation=3), # to maintain same ratio w.r.t. 224 images
|
||||
)
|
||||
t.append(transforms.CenterCrop(args.input_size))
|
||||
|
||||
t.append(transforms.ToTensor())
|
||||
t.append(transforms.Normalize(mean=IMAGENET_DEFAULT_MEAN, std=IMAGENET_DEFAULT_STD))
|
||||
return transforms.Compose(t)
|
||||
|
||||
|
||||
def get_sentencepiece_model_for_beit3(args):
|
||||
from transformers import XLMRobertaTokenizer
|
||||
return XLMRobertaTokenizer(args.sentencepiece_model)
|
||||
|
||||
|
||||
def create_dataset_by_split(args, split, is_train=True):
|
||||
transform = build_transform(is_train=is_train, args=args)
|
||||
dataset_class = task2dataset[args.task]
|
||||
tokenizer = get_sentencepiece_model_for_beit3(args)
|
||||
|
||||
opt_kwargs = {}
|
||||
if args.task in ["coco_captioning", "nocaps"]:
|
||||
opt_kwargs["mask_prob"] = args.captioning_mask_prob
|
||||
|
||||
dataset = dataset_class(
|
||||
data_path=args.data_path, split=split,
|
||||
transform=transform, tokenizer=tokenizer,
|
||||
num_max_bpe_tokens=args.num_max_bpe_tokens,
|
||||
task=args.task, **opt_kwargs,
|
||||
)
|
||||
if is_train:
|
||||
batch_size = args.batch_size
|
||||
elif hasattr(args, "eval_batch_size") and args.eval_batch_size is not None:
|
||||
batch_size = args.eval_batch_size
|
||||
else:
|
||||
batch_size = int(args.batch_size * 1.5)
|
||||
|
||||
return create_dataloader(
|
||||
dataset, is_train=is_train, batch_size=batch_size,
|
||||
num_workers=args.num_workers, pin_mem=args.pin_mem, dist_eval=args.dist_eval,
|
||||
)
|
||||
|
||||
|
||||
def create_downstream_dataset(args, is_eval=False):
|
||||
if is_eval:
|
||||
return create_dataset_by_split(args, split="test", is_train=False)
|
||||
else:
|
||||
return \
|
||||
create_dataset_by_split(args, split="train", is_train=True), \
|
||||
create_dataset_by_split(args, split="val", is_train=True)
|
||||
@@ -0,0 +1,598 @@
|
||||
# --------------------------------------------------------
|
||||
# Image as a Foreign Language: BEiT Pretraining for Vision and Vision-Language Tasks (https://arxiv.org/abs/2208.10442)
|
||||
# Github source: https://github.com/microsoft/unilm/tree/master/beit3
|
||||
# Copyright (c) 2023 Microsoft
|
||||
# Licensed under The MIT License [see LICENSE for details]
|
||||
# --------------------------------------------------------'
|
||||
|
||||
import math
|
||||
import sys
|
||||
import json
|
||||
from typing import Iterable, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from timm.utils import ModelEma
|
||||
from timm.utils import accuracy, ModelEma
|
||||
from timm.loss import LabelSmoothingCrossEntropy, SoftTargetCrossEntropy
|
||||
from datasets import get_sentencepiece_model_for_beit3
|
||||
|
||||
import utils
|
||||
|
||||
|
||||
class TaskHandler(object):
|
||||
def __init__(self) -> None:
|
||||
self.metric_logger = None
|
||||
self.split = None
|
||||
|
||||
def train_batch(self, model, **kwargs):
|
||||
raise NotImplementedError()
|
||||
|
||||
def eval_batch(self, model, **kwargs):
|
||||
raise NotImplementedError()
|
||||
|
||||
def before_eval(self, metric_logger, data_loader, **kwargs):
|
||||
self.metric_logger = metric_logger
|
||||
self.split = data_loader.dataset.split
|
||||
|
||||
def after_eval(self, **kwargs):
|
||||
raise NotImplementedError()
|
||||
|
||||
|
||||
class NLVR2Handler(TaskHandler):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.criterion = torch.nn.CrossEntropyLoss()
|
||||
|
||||
def train_batch(self, model, image, image2, language_tokens, padding_mask, label):
|
||||
logits = model(
|
||||
image_a=image, image_b=image2,
|
||||
text_description=language_tokens,
|
||||
padding_mask=padding_mask)
|
||||
acc = (logits.max(-1)[-1] == label).float().mean()
|
||||
return {
|
||||
"loss": self.criterion(input=logits, target=label),
|
||||
"acc": acc,
|
||||
}
|
||||
|
||||
def eval_batch(self, model, image, image2, language_tokens, padding_mask, label):
|
||||
logits = model(
|
||||
image_a=image, image_b=image2,
|
||||
text_description=language_tokens,
|
||||
padding_mask=padding_mask)
|
||||
batch_size = language_tokens.shape[0]
|
||||
acc = (logits.max(-1)[-1] == label).float().sum(0) * 100.0 / batch_size
|
||||
self.metric_logger.meters['acc'].update(acc.item(), n=batch_size)
|
||||
|
||||
def after_eval(self, **kwargs):
|
||||
print('* Acc {acc.global_avg:.3f}'.format(acc=self.metric_logger.acc))
|
||||
return {k: meter.global_avg for k, meter in self.metric_logger.meters.items()}, "acc"
|
||||
|
||||
|
||||
class ImageNetHandler(TaskHandler):
|
||||
def __init__(self, args) -> None:
|
||||
super().__init__()
|
||||
mixup_active = args.mixup > 0 or args.cutmix > 0. or args.cutmix_minmax is not None
|
||||
if mixup_active:
|
||||
# smoothing is handled with mixup label transform
|
||||
self.criterion = SoftTargetCrossEntropy()
|
||||
elif args.label_smoothing > 0.:
|
||||
self.criterion = LabelSmoothingCrossEntropy(smoothing=args.label_smoothing)
|
||||
else:
|
||||
self.criterion = torch.nn.CrossEntropyLoss()
|
||||
|
||||
def train_batch(self, model, image, label):
|
||||
logits = model(image=image)
|
||||
return {
|
||||
"loss": self.criterion(logits, label),
|
||||
}
|
||||
|
||||
def eval_batch(self, model, image, label):
|
||||
logits = model(image=image)
|
||||
batch_size = image.shape[0]
|
||||
acc1, acc5 = accuracy(logits, label, topk=(1, 5))
|
||||
self.metric_logger.meters['acc1'].update(acc1.item(), n=batch_size)
|
||||
self.metric_logger.meters['acc5'].update(acc5.item(), n=batch_size)
|
||||
|
||||
def after_eval(self, **kwargs):
|
||||
print('* Acc@1 {top1.global_avg:.3f} Acc@5 {top5.global_avg:.3f}'
|
||||
.format(top1=self.metric_logger.acc1, top5=self.metric_logger.acc5))
|
||||
return {k: meter.global_avg for k, meter in self.metric_logger.meters.items()}, "acc1"
|
||||
|
||||
|
||||
class RetrievalHandler(TaskHandler):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.image_feats = []
|
||||
self.text_feats = []
|
||||
self.image_ids = []
|
||||
self.metric_logger = None
|
||||
|
||||
def train_batch(self, model, image, language_tokens, padding_mask, image_id):
|
||||
loss, vision_cls, language_cls = model(
|
||||
image=image, text_description=language_tokens, padding_mask=padding_mask)
|
||||
return {
|
||||
"loss": loss,
|
||||
}
|
||||
|
||||
def before_eval(self, metric_logger, **kwargs):
|
||||
self.image_feats.clear()
|
||||
self.text_feats.clear()
|
||||
self.image_ids.clear()
|
||||
self.metric_logger = metric_logger
|
||||
|
||||
def eval_batch(self, model, image, language_tokens, padding_mask, image_id):
|
||||
vision_cls, _ = model(image=image, only_infer=True)
|
||||
_, language_cls = model(
|
||||
text_description=language_tokens, padding_mask=padding_mask, only_infer=True)
|
||||
|
||||
self.image_feats.append(vision_cls.clone())
|
||||
self.text_feats.append(language_cls.clone())
|
||||
self.image_ids.append(image_id.clone())
|
||||
|
||||
def after_eval(self, **kwargs):
|
||||
image_feats = {}
|
||||
for feats, ids in zip(self.image_feats, self.image_ids):
|
||||
for i, _idx in enumerate(ids):
|
||||
idx = _idx.item()
|
||||
if idx not in image_feats:
|
||||
image_feats[idx] = feats[i]
|
||||
|
||||
tiids = torch.cat(self.image_ids, dim=0)
|
||||
iids = []
|
||||
sorted_tensors = []
|
||||
for key in sorted(image_feats.keys()):
|
||||
sorted_tensors.append(image_feats[key].view(1, -1))
|
||||
iids.append(key)
|
||||
|
||||
image_cls_feats = torch.cat(sorted_tensors, dim=0)
|
||||
text_cls_feats = torch.cat(self.text_feats, dim=0)
|
||||
|
||||
scores = image_cls_feats @ text_cls_feats.t()
|
||||
iids = torch.LongTensor(iids).to(scores.device)
|
||||
|
||||
print("scores: {}".format(scores.size()))
|
||||
print("iids: {}".format(iids.size()))
|
||||
print("tiids: {}".format(tiids.size()))
|
||||
|
||||
topk10 = scores.topk(10, dim=1)
|
||||
topk5 = scores.topk(5, dim=1)
|
||||
topk1 = scores.topk(1, dim=1)
|
||||
|
||||
topk10_iids = tiids[topk10.indices]
|
||||
topk5_iids = tiids[topk5.indices]
|
||||
topk1_iids = tiids[topk1.indices]
|
||||
|
||||
tr_r10 = (iids.unsqueeze(1) == topk10_iids).float().max(dim=1)[0].mean()
|
||||
tr_r5 = (iids.unsqueeze(1) == topk5_iids).float().max(dim=1)[0].mean()
|
||||
tr_r1 = (iids.unsqueeze(1) == topk1_iids).float().max(dim=1)[0].mean()
|
||||
|
||||
topk10 = scores.topk(10, dim=0)
|
||||
topk5 = scores.topk(5, dim=0)
|
||||
topk1 = scores.topk(1, dim=0)
|
||||
topk10_iids = iids[topk10.indices]
|
||||
topk5_iids = iids[topk5.indices]
|
||||
topk1_iids = iids[topk1.indices]
|
||||
|
||||
ir_r10 = (tiids.unsqueeze(0) == topk10_iids).float().max(dim=0)[0].mean()
|
||||
ir_r5 = (tiids.unsqueeze(0) == topk5_iids).float().max(dim=0)[0].mean()
|
||||
ir_r1 = (tiids.unsqueeze(0) == topk1_iids).float().max(dim=0)[0].mean()
|
||||
|
||||
eval_result = {
|
||||
"tr_r10": tr_r10.item() * 100.0,
|
||||
"tr_r5": tr_r5.item() * 100.0,
|
||||
"tr_r1": tr_r1.item() * 100.0,
|
||||
"ir_r10": ir_r10.item() * 100.0,
|
||||
"ir_r5": ir_r5.item() * 100.0,
|
||||
"ir_r1": ir_r1.item() * 100.0,
|
||||
"average_score": 100.0 * (tr_r1 + tr_r5 + tr_r10 + ir_r1 + ir_r5 + ir_r10).item() / 6.0,
|
||||
}
|
||||
|
||||
print('* Eval result = %s' % json.dumps(eval_result))
|
||||
return eval_result, "average_score"
|
||||
|
||||
|
||||
class VQAHandler(TaskHandler):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.predictions = []
|
||||
self.criterion = nn.BCEWithLogitsLoss(reduction='mean')
|
||||
self.label2ans = None
|
||||
|
||||
def train_batch(self, model, image, language_tokens, padding_mask, labels):
|
||||
logits = model(
|
||||
image=image, question=language_tokens,
|
||||
padding_mask=padding_mask)
|
||||
return {
|
||||
"loss": self.criterion(input=logits.float(), target=labels.float()) * labels.shape[1],
|
||||
}
|
||||
|
||||
def before_eval(self, metric_logger, data_loader, **kwargs):
|
||||
self.predictions.clear()
|
||||
self.metric_logger = metric_logger
|
||||
self.label2ans = data_loader.dataset.label2ans
|
||||
|
||||
def eval_batch(self, model, image, language_tokens, padding_mask, labels=None, qid=None):
|
||||
logits = model(
|
||||
image=image, question=language_tokens,
|
||||
padding_mask=padding_mask)
|
||||
batch_size = language_tokens.shape[0]
|
||||
if labels is not None:
|
||||
scores = utils.VQAScore()(logits, labels) * 100.0
|
||||
self.metric_logger.meters['score'].update(scores.item(), n=batch_size)
|
||||
else:
|
||||
_, preds = logits.max(-1)
|
||||
for image_id, pred in zip(qid, preds):
|
||||
self.predictions.append({
|
||||
"question_id": image_id.item(),
|
||||
"answer": self.label2ans[pred.item()],
|
||||
})
|
||||
|
||||
def after_eval(self, **kwargs):
|
||||
if len(self.predictions) == 0:
|
||||
print('* Score {score.global_avg:.3f}'.format(score=self.metric_logger.score))
|
||||
return {k: meter.global_avg for k, meter in self.metric_logger.meters.items()}, "score"
|
||||
else:
|
||||
return self.predictions, "prediction"
|
||||
|
||||
|
||||
class CaptioningHandler(TaskHandler):
|
||||
def __init__(self, args) -> None:
|
||||
super().__init__()
|
||||
self.predictions = []
|
||||
self.criterion = utils.BertCaptioningLoss(args.label_smoothing, args.drop_worst_ratio, args.drop_worst_after)
|
||||
self.tokenizer = get_sentencepiece_model_for_beit3(args)
|
||||
self.num_beams = args.num_beams
|
||||
self.max_len = args.num_max_bpe_tokens
|
||||
self.length_penalty = args.length_penalty
|
||||
self.vocab_size = args.vocab_size
|
||||
|
||||
def train_batch(self, model, image, language_tokens, masked_tokens, language_masked_pos, padding_mask, image_id, global_step):
|
||||
logits, _ = model(
|
||||
image=image, text_ids=masked_tokens, padding_mask=padding_mask, language_masked_pos=language_masked_pos, image_id=image_id)
|
||||
masked_labels = language_tokens[language_masked_pos.bool()]
|
||||
score = torch.max(logits, -1)[1].data == masked_labels
|
||||
acc = torch.sum(score.float()) / torch.sum(language_masked_pos)
|
||||
return {
|
||||
"loss": self.criterion(logits, masked_labels, global_step),
|
||||
"acc": acc
|
||||
}
|
||||
|
||||
def before_eval(self, metric_logger, data_loader, **kwargs):
|
||||
self.predictions.clear()
|
||||
self.metric_logger = metric_logger
|
||||
|
||||
def eval_batch(self, model, image, image_id=None):
|
||||
cur_len = 2
|
||||
num_keep_best = 1
|
||||
TOPN_PER_BEAM = 3
|
||||
|
||||
batch_size = image.size(0)
|
||||
mask_id = self.tokenizer.mask_token_id
|
||||
cls_id = self.tokenizer.cls_token_id
|
||||
pad_id = self.tokenizer.pad_token_id
|
||||
sep_id = self.tokenizer.sep_token_id
|
||||
eos_token_ids = [sep_id]
|
||||
|
||||
cls_ids = torch.full(
|
||||
(batch_size, 1), cls_id, dtype=torch.long, device=image.device
|
||||
)
|
||||
mask_ids = torch.full(
|
||||
(batch_size, 1), mask_id, dtype=torch.long, device=image.device
|
||||
)
|
||||
cur_input_ids = torch.cat([cls_ids, mask_ids], dim=1)
|
||||
tmp_ids = torch.full(
|
||||
(batch_size, self.max_len-1), mask_id, dtype=torch.long, device=image.device
|
||||
)
|
||||
decoding_results = torch.cat([cls_ids, tmp_ids], dim=1)
|
||||
|
||||
# Expand input to num beams
|
||||
cur_input_ids = cur_input_ids.unsqueeze(1).expand(batch_size, self.num_beams, cur_len)
|
||||
cur_input_ids = cur_input_ids.contiguous().view(batch_size * self.num_beams, cur_len) # (batch_size * num_beams, cur_len)
|
||||
decoding_results = decoding_results.unsqueeze(1).expand(batch_size, self.num_beams, self.max_len)
|
||||
decoding_results = decoding_results.contiguous().view(batch_size * self.num_beams, self.max_len) # (batch_size * num_beams, cur_len)
|
||||
image = image.unsqueeze(1).expand(batch_size, self.num_beams, image.size(-3), image.size(-2), image.size(-1))
|
||||
image = image.contiguous().view(batch_size * self.num_beams, image.size(-3), image.size(-2), image.size(-1))
|
||||
|
||||
generated_hyps = [
|
||||
utils.BeamHypotheses(
|
||||
num_keep_best, self.max_len, length_penalty=self.length_penalty, early_stopping=False
|
||||
) for _ in range(batch_size)
|
||||
]
|
||||
# scores for each sentence in the beam
|
||||
beam_scores = torch.zeros((batch_size, self.num_beams), dtype=torch.float, device=cur_input_ids.device)
|
||||
beam_scores[:, 1:] = -1e9
|
||||
beam_scores = beam_scores.view(-1) # shape (batch_size * num_beams,)
|
||||
|
||||
# done sentences
|
||||
done = [False for _ in range(batch_size)]
|
||||
incremental_state = {}
|
||||
|
||||
while cur_len <= self.max_len:
|
||||
next_token_idx = 1
|
||||
padding_masks = torch.full(
|
||||
cur_input_ids.shape, 0, dtype=torch.long, device=image.device
|
||||
)
|
||||
input_image = image
|
||||
if cur_len != 2:
|
||||
input_image = None
|
||||
|
||||
outputs, incremental_state_next = model(
|
||||
image=input_image, text_ids=cur_input_ids, language_masked_pos=None,
|
||||
padding_mask=padding_masks, text_len=cur_len, incremental_state=incremental_state)
|
||||
incremental_state = incremental_state_next
|
||||
|
||||
# assert outputs.shape[1] == token_len
|
||||
scores = outputs[:, next_token_idx, :] # (batch_size * num_beams, vocab_size)
|
||||
scores = F.log_softmax(scores, dim=-1) # (batch_size * num_beams, vocab_size)
|
||||
assert scores.size() == (batch_size * self.num_beams, self.vocab_size)
|
||||
# Add the log prob of the new beams to the log prob of the beginning of the sequence (sum of logs == log of the product)
|
||||
_scores = scores + beam_scores[:, None].expand_as(scores) # (batch_size * num_beams, vocab_size)
|
||||
# re-organize to group the beam together (we are keeping top hypothesis accross beams)
|
||||
_scores = _scores.view(batch_size, self.num_beams * self.vocab_size) # (batch_size, num_beams * vocab_size)
|
||||
next_scores, next_words = torch.topk(_scores, TOPN_PER_BEAM * self.num_beams, dim=1, largest=True, sorted=True)
|
||||
assert next_scores.size() == next_words.size() == (batch_size, TOPN_PER_BEAM * self.num_beams)
|
||||
|
||||
# next batch beam content
|
||||
# list of (batch_size * num_beams) tuple(next hypothesis score, next word, current position in the batch)
|
||||
next_batch_beam = []
|
||||
# for each sentence
|
||||
for batch_ex in range(batch_size):
|
||||
# if we are done with this sentence
|
||||
done[batch_ex] = done[batch_ex] or generated_hyps[batch_ex].is_done(next_scores[batch_ex].max().item())
|
||||
if done[batch_ex]:
|
||||
next_batch_beam.extend([(0, pad_id, 0)] * self.num_beams) # pad the batch
|
||||
continue
|
||||
|
||||
# next sentence beam content
|
||||
next_sent_beam = []
|
||||
for idx, score in zip(next_words[batch_ex], next_scores[batch_ex]):
|
||||
# get beam and word IDs
|
||||
beam_id = idx // self.vocab_size
|
||||
word_id = idx % self.vocab_size
|
||||
# end of sentence, or next word
|
||||
# if word_id.item() in eos_token_ids or cur_len + 1 == max_len:
|
||||
if (word_id.item() in eos_token_ids and cur_len + 1 <= self.max_len) or (cur_len + 1 == self.max_len):
|
||||
generated_hyps[batch_ex].add(
|
||||
decoding_results[batch_ex * self.num_beams + beam_id, :cur_len].clone(), score.item()
|
||||
)
|
||||
else:
|
||||
next_sent_beam.append((score, word_id, batch_ex * self.num_beams + beam_id))
|
||||
# the beam for next step is full
|
||||
if len(next_sent_beam) == self.num_beams:
|
||||
break
|
||||
|
||||
# update next beam content
|
||||
if cur_len + 1 == self.max_len:
|
||||
assert len(next_sent_beam) == 0
|
||||
else:
|
||||
assert len(next_sent_beam) == self.num_beams
|
||||
|
||||
if len(next_sent_beam) == 0:
|
||||
next_sent_beam = [(0, pad_id, 0)] * self.num_beams # pad the batch
|
||||
next_batch_beam.extend(next_sent_beam)
|
||||
assert len(next_batch_beam) == self.num_beams * (batch_ex + 1)
|
||||
|
||||
# sanity check / prepare next batch
|
||||
assert len(next_batch_beam) == batch_size * self.num_beams
|
||||
beam_scores = beam_scores.new([x[0] for x in next_batch_beam])
|
||||
beam_words = cur_input_ids.new([x[1] for x in next_batch_beam])
|
||||
beam_idx = cur_input_ids.new([x[2] for x in next_batch_beam])
|
||||
|
||||
# re-order batch
|
||||
cur_input_ids = cur_input_ids[beam_idx, :]
|
||||
decoding_results = decoding_results[beam_idx, :]
|
||||
for module in incremental_state:
|
||||
for key in incremental_state[module]:
|
||||
result = incremental_state[module][key].index_select(0, beam_idx)
|
||||
incremental_state[module][key] = result[:,:,:-1,:]
|
||||
|
||||
next_ids = torch.full(
|
||||
(batch_size * self.num_beams, 1), mask_id, dtype=torch.long, device=image.device
|
||||
)
|
||||
cur_input_ids = torch.cat([beam_words.unsqueeze(1), next_ids], dim=1)
|
||||
decoding_results[:, cur_len-1] = beam_words
|
||||
# update current length
|
||||
cur_len = cur_len + 1
|
||||
# stop when we are done with each sentence
|
||||
if all(done):
|
||||
break
|
||||
|
||||
# select the best hypotheses
|
||||
tgt_len = torch.ones(batch_size, num_keep_best, dtype=torch.long)
|
||||
logprobs = torch.zeros(batch_size, num_keep_best,
|
||||
dtype=torch.float).fill_(-1e5).to(cur_input_ids.device)
|
||||
all_best = []
|
||||
|
||||
for i, hypotheses in enumerate(generated_hyps):
|
||||
best = []
|
||||
hyp_scores = torch.tensor([x[0] for x in hypotheses.hyp])
|
||||
_, best_indices = torch.topk(hyp_scores,
|
||||
min(num_keep_best, len(hyp_scores)), largest=True)
|
||||
for best_idx, hyp_idx in enumerate(best_indices):
|
||||
conf, best_hyp = hypotheses.hyp[hyp_idx]
|
||||
best.append(best_hyp)
|
||||
logprobs[i, best_idx] = conf
|
||||
tgt_len[i, best_idx] = len(best_hyp) + 1 # +1 for the <EOS> symbol
|
||||
all_best.append(best)
|
||||
|
||||
# generate target batch, pad to the same length
|
||||
decoded = cur_input_ids.new(batch_size, num_keep_best, self.max_len).fill_(pad_id)
|
||||
for batch_idx, best in enumerate(all_best):
|
||||
for best_idx, hypo in enumerate(best):
|
||||
decoded[batch_idx, best_idx, : tgt_len[batch_idx, best_idx] - 1] = hypo
|
||||
decoded[batch_idx, best_idx, tgt_len[batch_idx, best_idx] - 1] = eos_token_ids[0]
|
||||
|
||||
captions = self.tokenizer.batch_decode(decoded.squeeze(1), skip_special_tokens=True)
|
||||
for qid, pred in zip(image_id, captions):
|
||||
self.predictions.append({
|
||||
"image_id": qid.item(),
|
||||
"caption": pred,
|
||||
})
|
||||
|
||||
def after_eval(self, **kwargs):
|
||||
return self.predictions, "prediction"
|
||||
|
||||
|
||||
def get_handler(args):
|
||||
if args.task == "nlvr2":
|
||||
return NLVR2Handler()
|
||||
elif args.task == "vqav2":
|
||||
return VQAHandler()
|
||||
elif args.task in ("flickr30k", "coco_retrieval"):
|
||||
return RetrievalHandler()
|
||||
elif args.task in ("coco_captioning", "nocaps"):
|
||||
return CaptioningHandler(args)
|
||||
elif args.task in ("imagenet"):
|
||||
return ImageNetHandler(args)
|
||||
else:
|
||||
raise NotImplementedError("Sorry, %s is not support." % args.task)
|
||||
|
||||
|
||||
def train_one_epoch(
|
||||
model: torch.nn.Module, data_loader: Iterable,
|
||||
optimizer: torch.optim.Optimizer, device: torch.device,
|
||||
handler: TaskHandler, epoch: int, start_steps: int,
|
||||
lr_schedule_values: list, loss_scaler, max_norm: float = 0,
|
||||
update_freq: int = 1, model_ema: Optional[ModelEma] = None,
|
||||
log_writer: Optional[utils.TensorboardLogger] = None,
|
||||
task = None, mixup_fn=None,
|
||||
):
|
||||
model.train(True)
|
||||
metric_logger = utils.MetricLogger(delimiter=" ")
|
||||
metric_logger.add_meter('lr', utils.SmoothedValue(window_size=1, fmt='{value:.6f}'))
|
||||
metric_logger.add_meter('min_lr', utils.SmoothedValue(window_size=1, fmt='{value:.6f}'))
|
||||
header = 'Epoch: [{}]'.format(epoch)
|
||||
print_freq = 10
|
||||
|
||||
if loss_scaler is None:
|
||||
model.zero_grad()
|
||||
model.micro_steps = 0
|
||||
else:
|
||||
optimizer.zero_grad()
|
||||
|
||||
for data_iter_step, data in enumerate(metric_logger.log_every(data_loader, print_freq, header)):
|
||||
step = data_iter_step // update_freq
|
||||
global_step = start_steps + step # global training iteration
|
||||
# Update LR & WD for the first acc
|
||||
if lr_schedule_values is not None and data_iter_step % update_freq == 0:
|
||||
for i, param_group in enumerate(optimizer.param_groups):
|
||||
if lr_schedule_values is not None:
|
||||
param_group["lr"] = lr_schedule_values[global_step] * param_group["lr_scale"]
|
||||
# put input data into cuda
|
||||
for tensor_key in data.keys():
|
||||
data[tensor_key] = data[tensor_key].to(device, non_blocking=True)
|
||||
# print("input %s = %s" % (tensor_key, data[tensor_key]))
|
||||
if loss_scaler is None and tensor_key.startswith("image"):
|
||||
data[tensor_key] = data[tensor_key].half()
|
||||
|
||||
# mixup for imagenet finetuning
|
||||
if mixup_fn is not None:
|
||||
data["image"], data["label"] = mixup_fn(data["image"], data["label"])
|
||||
|
||||
if task in ["coco_captioning", "nocaps"]:
|
||||
data["global_step"] = global_step
|
||||
|
||||
if loss_scaler is None:
|
||||
results = handler.train_batch(model, **data)
|
||||
else:
|
||||
with torch.cuda.amp.autocast():
|
||||
results = handler.train_batch(model, **data)
|
||||
|
||||
loss = results.pop("loss")
|
||||
loss_value = loss.item()
|
||||
|
||||
if not math.isfinite(loss_value):
|
||||
print("Loss is {}, stopping training".format(loss_value))
|
||||
sys.exit(1)
|
||||
|
||||
if loss_scaler is None:
|
||||
loss /= update_freq
|
||||
model.backward(loss)
|
||||
model.step()
|
||||
|
||||
if (data_iter_step + 1) % update_freq == 0:
|
||||
# model.zero_grad()
|
||||
# Deepspeed will call step() & model.zero_grad() automatic
|
||||
if model_ema is not None:
|
||||
model_ema.update(model)
|
||||
grad_norm = None
|
||||
loss_scale_value = utils.get_loss_scale_for_deepspeed(model)
|
||||
else:
|
||||
# this attribute is added by timm on one optimizer (adahessian)
|
||||
is_second_order = hasattr(optimizer, 'is_second_order') and optimizer.is_second_order
|
||||
loss /= update_freq
|
||||
grad_norm = loss_scaler(loss, optimizer, clip_grad=max_norm,
|
||||
parameters=model.parameters(), create_graph=is_second_order,
|
||||
update_grad=(data_iter_step + 1) % update_freq == 0)
|
||||
if (data_iter_step + 1) % update_freq == 0:
|
||||
optimizer.zero_grad()
|
||||
if model_ema is not None:
|
||||
model_ema.update(model)
|
||||
loss_scale_value = loss_scaler.state_dict()["scale"]
|
||||
|
||||
torch.cuda.synchronize()
|
||||
|
||||
metric_logger.update(loss=loss_value)
|
||||
metric_logger.update(loss_scale=loss_scale_value)
|
||||
min_lr = 10.
|
||||
max_lr = 0.
|
||||
for group in optimizer.param_groups:
|
||||
min_lr = min(min_lr, group["lr"])
|
||||
max_lr = max(max_lr, group["lr"])
|
||||
|
||||
metric_logger.update(lr=max_lr)
|
||||
metric_logger.update(min_lr=min_lr)
|
||||
weight_decay_value = None
|
||||
for group in optimizer.param_groups:
|
||||
if group["weight_decay"] > 0:
|
||||
weight_decay_value = group["weight_decay"]
|
||||
metric_logger.update(weight_decay=weight_decay_value)
|
||||
metric_logger.update(grad_norm=grad_norm)
|
||||
|
||||
if log_writer is not None:
|
||||
kwargs = {
|
||||
"loss": loss_value,
|
||||
}
|
||||
for key in results:
|
||||
kwargs[key] = results[key]
|
||||
log_writer.update(head="train", **kwargs)
|
||||
|
||||
kwargs = {
|
||||
"loss_scale": loss_scale_value,
|
||||
"lr": max_lr,
|
||||
"min_lr": min_lr,
|
||||
"weight_decay": weight_decay_value,
|
||||
"grad_norm": grad_norm,
|
||||
}
|
||||
log_writer.update(head="opt", **kwargs)
|
||||
log_writer.set_step()
|
||||
|
||||
# gather the stats from all processes
|
||||
metric_logger.synchronize_between_processes()
|
||||
print("Averaged stats:", metric_logger)
|
||||
return {k: meter.global_avg for k, meter in metric_logger.meters.items()}
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def evaluate(data_loader, model, device, handler):
|
||||
metric_logger = utils.MetricLogger(delimiter=" ")
|
||||
header = 'Test:'
|
||||
|
||||
# switch to evaluation mode
|
||||
model.eval()
|
||||
handler.before_eval(metric_logger=metric_logger, data_loader=data_loader)
|
||||
|
||||
for data in metric_logger.log_every(data_loader, 10, header):
|
||||
for tensor_key in data.keys():
|
||||
data[tensor_key] = data[tensor_key].to(device, non_blocking=True)
|
||||
|
||||
with torch.cuda.amp.autocast():
|
||||
handler.eval_batch(model=model, **data)
|
||||
|
||||
# gather the stats from all processes
|
||||
metric_logger.synchronize_between_processes()
|
||||
|
||||
return handler.after_eval()
|
||||
@@ -0,0 +1,176 @@
|
||||
# Fine-tuning BEiT-3 on Image Captioning
|
||||
|
||||
## COCO Captioning Setup
|
||||
|
||||
1. [Setup environment](../README.md#setup).
|
||||
2. Download [2014 train images](http://images.cocodataset.org/zips/train2014.zip), [2014 val images](http://images.cocodataset.org/zips/val2014.zip) and [karpathy split](https://cs.stanford.edu/people/karpathy/deepimagesent/caption_datasets.zip), then organize the dataset as following structure:
|
||||
|
||||
```
|
||||
/path/to/your_data/
|
||||
train2014/
|
||||
COCO_train2014_000000000009.jpg
|
||||
...
|
||||
val2014/
|
||||
COCO_val2014_000000000042.jpg
|
||||
...
|
||||
dataset_coco.json
|
||||
```
|
||||
|
||||
We then generate the index json files using the following command. [beit3.spm](https://github.com/addf400/files/releases/download/beit3/beit3.spm) is the sentencepiece model used for tokenizing texts.
|
||||
```
|
||||
from datasets import CaptioningDataset
|
||||
from transformers import XLMRobertaTokenizer
|
||||
|
||||
tokenizer = XLMRobertaTokenizer("/your_beit3_model_path/beit3.spm")
|
||||
|
||||
CaptioningDataset.make_coco_captioning_dataset_index(
|
||||
data_path="/path/to/your_data",
|
||||
tokenizer=tokenizer,
|
||||
)
|
||||
```
|
||||
|
||||
|
||||
## NoCaps Setup
|
||||
|
||||
1. [Setup environment](README.md#setup).
|
||||
2. Download [NoCaps val set](https://nocaps.s3.amazonaws.com/nocaps_val_4500_captions.json), [NoCaps test set](https://s3.amazonaws.com/nocaps/nocaps_test_image_info.json) and download imags using the urls in val and test json files, then organize the dataset as following structure:
|
||||
|
||||
```
|
||||
/path/to/your_data/
|
||||
val/
|
||||
09c863d76bcf6b00.jpg
|
||||
...
|
||||
test/
|
||||
19dc6913830a0a21.jpg
|
||||
...
|
||||
nocaps_val_4500_captions.json
|
||||
nocaps_test_image_info.json
|
||||
```
|
||||
|
||||
We then generate the index json files using the following command. [beit3.spm](https://github.com/addf400/files/releases/download/beit3/beit3.spm) is the sentencepiece model used for tokenizing texts.
|
||||
```
|
||||
from datasets import CaptioningDataset
|
||||
from transformers import XLMRobertaTokenizer
|
||||
|
||||
tokenizer = XLMRobertaTokenizer("/your_beit3_model_path/beit3.spm")
|
||||
|
||||
CaptioningDataset.make_nocaps_captioning_dataset_index(
|
||||
data_path="/path/to/your_data",
|
||||
)
|
||||
```
|
||||
We use COCO captioning training set as the training data of NoCaps.
|
||||
|
||||
|
||||
## Example: Fine-tuning BEiT-3 on Captioning
|
||||
|
||||
The BEiT-3 **base** model can be fine-tuned on captioning tasks using 8 V100-32GB:
|
||||
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
|
||||
--model beit3_base_patch16_480 \
|
||||
--input_size 480 \
|
||||
--task coco_captioning \
|
||||
--batch_size 32 \
|
||||
--layer_decay 1.0 \
|
||||
--lr 4e-5 \
|
||||
--randaug \
|
||||
--epochs 10 \
|
||||
--warmup_epochs 1 \
|
||||
--drop_path 0.1 \
|
||||
--sentencepiece_model /your_beit3_model_path/beit3.spm \
|
||||
--finetune /your_beit3_model_path/beit3_base_patch16_224.pth \
|
||||
--data_path /path/to/your_data \
|
||||
--output_dir /path/to/save/your_model \
|
||||
--log_dir /path/to/save/your_model/log \
|
||||
--weight_decay 0.05 \
|
||||
--seed 42 \
|
||||
--save_ckpt_freq 5 \
|
||||
--num_max_bpe_tokens 32 \
|
||||
--captioning_mask_prob 0.7 \
|
||||
--drop_worst_after 12000 \
|
||||
--dist_eval \
|
||||
--checkpoint_activations \
|
||||
--enable_deepspeed
|
||||
```
|
||||
- `--batch_size`: batch size per GPU. Effective batch size = `number of GPUs` * `--batch_size` * `--update_freq`. So in the above example, the effective batch size is `8*32 = 256`.
|
||||
- `--finetune`: weight path of your pretrained models; please download the pretrained model weights in [README.md](../README.md#pretrained-models).
|
||||
- `--task`: **coco_captioning** for COCO captioning and **nocaps** for NoCaps dataset.
|
||||
- `lr`: 4e-5 for COCO captioning and 1e-5 for NoCaps.
|
||||
- `--enable_deepspeed`: optional. If you use apex, please enable deepspeed.
|
||||
- `--checkpoint_activations`: using gradient checkpointing for saving GPU memory.
|
||||
|
||||
|
||||
The BEiT-3 **large** model can be fine-tuned on captioning tasks using 8 V100-32GB:
|
||||
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
|
||||
--model beit3_large_patch16_480 \
|
||||
--input_size 480 \
|
||||
--task coco_captioning \
|
||||
--batch_size 32 \
|
||||
--layer_decay 1.0 \
|
||||
--lr 8e-6 \
|
||||
--randaug \
|
||||
--epochs 10 \
|
||||
--warmup_epochs 1 \
|
||||
--drop_path 0.1 \
|
||||
--sentencepiece_model /your_beit3_model_path/beit3.spm \
|
||||
--finetune /your_beit3_model_path/beit3_large_patch16_224.pth \
|
||||
--data_path /path/to/your_data \
|
||||
--output_dir /path/to/save/your_model \
|
||||
--log_dir /path/to/save/your_model/log \
|
||||
--weight_decay 0.05 \
|
||||
--seed 42 \
|
||||
--save_ckpt_freq 5 \
|
||||
--num_max_bpe_tokens 32 \
|
||||
--captioning_mask_prob 0.7 \
|
||||
--drop_worst_after 12000 \
|
||||
--dist_eval \
|
||||
--checkpoint_activations \
|
||||
--enable_deepspeed
|
||||
```
|
||||
- `--batch_size`: batch size per GPU. Effective batch size = `number of GPUs` * `--batch_size` * `--update_freq`. So in the above example, the effective batch size is `8*32 = 256`.
|
||||
- `--finetune`: weight path of your pretrained models; please download the pretrained model weights in [README.md](../README.md#pretrained-models).
|
||||
- `--task`: **coco_captioning** for COCO captioning and **nocaps** for NoCaps dataset.
|
||||
- `lr`: 8e-6 for COCO captioning and NoCaps.
|
||||
- `--enable_deepspeed`: optional. If you use apex, please enable deepspeed.
|
||||
- `--checkpoint_activations`: using gradient checkpointing for saving GPU memory.
|
||||
|
||||
|
||||
## Example: Evaluate BEiT-3 Fine-tuned model on Captioning
|
||||
|
||||
- Get the prediction file of the fine-tuned BEiT3-base model on captioning with 8 V100-32GB:
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
|
||||
--model beit3_base_patch16_480 \
|
||||
--input_size 480 \
|
||||
--task coco_captioning \
|
||||
--batch_size 16 \
|
||||
--sentencepiece_model /your_beit3_model_path/beit3.spm \
|
||||
--finetune /your_beit3_model_path/beit3_base_patch16_480_coco_captioning.pth \
|
||||
--data_path /path/to/your_data \
|
||||
--output_dir /path/to/save/your_prediction \
|
||||
--eval \
|
||||
--dist_eval
|
||||
```
|
||||
- `--task`: **coco_captioning** for COCO captioning and **nocaps** for NoCaps dataset.
|
||||
- `--finetune`: **beit3_base_patch16_480_coco_captioning.pth** for COCO captioning and **beit3_base_patch16_480_nocaps.pth** for NoCaps dataset.
|
||||
|
||||
- Get the prediction file of the fine-tuned BEiT3-large model on captioning with 8 V100-32GB:
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
|
||||
--model beit3_large_patch16_480 \
|
||||
--input_size 480 \
|
||||
--task coco_captioning \
|
||||
--batch_size 16 \
|
||||
--sentencepiece_model /your_beit3_model_path/beit3.spm \
|
||||
--finetune /your_beit3_model_path/beit3_large_patch16_480_coco_captioning.pth \
|
||||
--data_path /path/to/your_data \
|
||||
--output_dir /path/to/save/your_prediction \
|
||||
--eval \
|
||||
--dist_eval
|
||||
```
|
||||
- `--task`: **coco_captioning** for COCO captioning and **nocaps** for NoCaps dataset.
|
||||
- `--finetune`: **beit3_large_patch16_480_coco_captioning.pth** for COCO captioning and **beit3_large_patch16_480_nocaps.pth** for NoCaps dataset.
|
||||
|
||||
Please then submit the prediction file in the `output_dir` to the [evaluation server](https://eval.ai/web/challenges/challenge-page/355/overview) to obtain the NoCaps val and test results.
|
||||
@@ -0,0 +1,138 @@
|
||||
# Fine-tuning BEiT-3 on ImageNet-1k (Image Classification)
|
||||
|
||||
|
||||
## Setup
|
||||
|
||||
1. [Setup environment](../README.md#setup).
|
||||
2. Download and extract ImageNet-1k from http://image-net.org/.
|
||||
|
||||
The directory structure is the standard layout of torchvision's [`datasets.ImageFolder`](https://pytorch.org/docs/stable/torchvision/datasets.html#imagefolder). The training and validation data are expected to be in the `train/` folder and `val/` folder, respectively:
|
||||
|
||||
```
|
||||
/path/to/imagenet/
|
||||
train/
|
||||
class1/
|
||||
img1.jpeg
|
||||
class2/
|
||||
img2.jpeg
|
||||
val/
|
||||
class1/
|
||||
img3.jpeg
|
||||
class/2
|
||||
img4.jpeg
|
||||
```
|
||||
|
||||
We then generate the index json files using the following command. [beit3.spm](https://github.com/addf400/files/releases/download/beit3/beit3.spm) is the sentencepiece model used for tokenizing texts.
|
||||
```
|
||||
from datasets import ImageNetDataset
|
||||
|
||||
ImageNetDataset.make_dataset_index(
|
||||
train_data_path = "/path/to/your_data/train",
|
||||
val_data_path = "/path/to/your_data/val",
|
||||
index_path = "/path/to/your_data"
|
||||
)
|
||||
```
|
||||
|
||||
|
||||
## Example: Fine-tuning BEiT-3 on ImageNet-1k (Image Classification)
|
||||
|
||||
The BEiT-3 **base** model can be finetuned on ImageNet-1k using 8 V100-32GB:
|
||||
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
|
||||
--model beit3_base_patch16_224 \
|
||||
--task imagenet \
|
||||
--batch_size 128 \
|
||||
--layer_decay 0.65 \
|
||||
--lr 7e-4 \
|
||||
--update_freq 1 \
|
||||
--epochs 50 \
|
||||
--warmup_epochs 5 \
|
||||
--drop_path 0.15 \
|
||||
--sentencepiece_model /your_beit3_model_path/beit3.spm \
|
||||
--finetune /your_beit3_model_path/beit3_base_patch16_224.pth \
|
||||
--data_path /path/to/your_data \
|
||||
--output_dir /path/to/save/your_model \
|
||||
--log_dir /path/to/save/your_model/log \
|
||||
--weight_decay 0.05 \
|
||||
--seed 42 \
|
||||
--save_ckpt_freq 5 \
|
||||
--dist_eval \
|
||||
--mixup 0.8 \
|
||||
--cutmix 1.0 \
|
||||
--enable_deepspeed
|
||||
```
|
||||
- `--batch_size`: batch size per GPU. Effective batch size = `number of GPUs` * `--batch_size` * `--update_freq`. So in the above example, the effective batch size is `8*128*1 = 1024`.
|
||||
- `--finetune`: weight path of your pretrained models; please download the pretrained model weights in [README.md](../README.md#pretrained-models)
|
||||
- `--enable_deepspeed`: optional. If you use apex, please enable deepspeed.
|
||||
|
||||
|
||||
The BEiT-3 **large** model can be finetuned on ImageNet-1k using a DGX box (8 V100-32GB):
|
||||
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
|
||||
--model beit3_large_patch16_224 \
|
||||
--task imagenet \
|
||||
--batch_size 128 \
|
||||
--layer_decay 0.8 \
|
||||
--lr 2e-4 \
|
||||
--update_freq 1 \
|
||||
--epochs 50 \
|
||||
--warmup_epochs 5 \
|
||||
--drop_path 0.25 \
|
||||
--sentencepiece_model /your_beit3_model_path/beit3.spm \
|
||||
--finetune /your_beit3_model_path/beit3_large_patch16_224.pth \
|
||||
--data_path /path/to/your_data \
|
||||
--output_dir /path/to/save/your_model \
|
||||
--log_dir /path/to/save/your_model/log \
|
||||
--weight_decay 0.05 \
|
||||
--seed 42 \
|
||||
--save_ckpt_freq 5 \
|
||||
--dist_eval \
|
||||
--mixup 0.8 \
|
||||
--cutmix 1.0 \
|
||||
--enable_deepspeed \
|
||||
--checkpoint_activations
|
||||
```
|
||||
- `--batch_size`: batch size per GPU. Effective batch size = `number of GPUs` * `--batch_size` * `--update_freq`. So in the above example, the effective batch size is `8*128 = 1024`.
|
||||
- `--finetune`: weight path of your pretrained models; please download the pretrained model weights in [README.md](../README.md#pretrained-models)
|
||||
- `--enable_deepspeed`: optional. If you use apex, please enable deepspeed.
|
||||
- `--checkpoint_activations`: using gradient checkpointing for saving GPU memory
|
||||
|
||||
## Example: Evaluate BEiT-3 Finetuned model on ImageNet-1k (Image Classification)
|
||||
|
||||
- Evaluate our fine-tuned BEiT3-base model on ImageNet val with a single GPU:
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=1 run_beit3_finetuning.py \
|
||||
--model beit3_base_patch16_224 \
|
||||
--task imagenet \
|
||||
--batch_size 128 \
|
||||
--sentencepiece_model /your_beit3_model_path/beit3.spm \
|
||||
--finetune /your_beit3_model_path/beit3_base_patch16_224_in1k.pth \
|
||||
--data_path /path/to/your_data \
|
||||
--eval \
|
||||
--dist_eval
|
||||
```
|
||||
|
||||
Expected results:
|
||||
```
|
||||
* Acc@1 85.400 Acc@5 97.630
|
||||
```
|
||||
|
||||
- Evaluate our fine-tuned BEiT3-large model on ImageNet val with a single GPU:
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=1 run_beit3_finetuning.py \
|
||||
--model beit3_large_patch16_224 \
|
||||
--task imagenet \
|
||||
--batch_size 128 \
|
||||
--sentencepiece_model /your_beit3_model_path/beit3.spm \
|
||||
--finetune /your_beit3_model_path/beit3_large_patch16_224_in1k.pth \
|
||||
--data_path /path/to/your_data \
|
||||
--eval \
|
||||
--dist_eval
|
||||
```
|
||||
|
||||
Expected results:
|
||||
```
|
||||
* Acc@1 87.580 Acc@5 98.326
|
||||
```
|
||||
@@ -0,0 +1,136 @@
|
||||
# Fine-tuning BEiT-3 on NLVR2 (Visual Reasoning)
|
||||
|
||||
|
||||
## Setup
|
||||
|
||||
1. [Setup environment](../README.md#setup).
|
||||
2. Clone the [repository](https://github.com/lil-lab/nlvr) and sign the [request form](https://goo.gl/forms/yS29stWnFWzrDBFH3) to download the images, then organize the dataset as following structure:
|
||||
|
||||
```
|
||||
/path/to/your_data/
|
||||
images/train/
|
||||
0/train-11670-0-img0.png
|
||||
...
|
||||
dev/
|
||||
dev-269-0-img0.png
|
||||
...
|
||||
test1/
|
||||
test1-261-0-img0.png
|
||||
...
|
||||
nlvr/ (nlvr repo)
|
||||
nlvr/
|
||||
nlvr2/
|
||||
```
|
||||
|
||||
We then generate the index json files using the following command. [beit3.spm](https://github.com/addf400/files/releases/download/beit3/beit3.spm) is the sentencepiece model used for tokenizing texts.
|
||||
```
|
||||
from datasets import NLVR2Dataset
|
||||
from transformers import XLMRobertaTokenizer
|
||||
|
||||
tokenizer = XLMRobertaTokenizer("/your_beit3_model_path/beit3.spm")
|
||||
|
||||
NLVR2Dataset.make_dataset_index(
|
||||
data_path="/path/to/your_data",
|
||||
tokenizer=tokenizer,
|
||||
nlvr_repo_path="/path/to/your_data/nlvr"
|
||||
)
|
||||
```
|
||||
|
||||
|
||||
## Example: Fine-tuning BEiT-3 on NLVR2 (Visual Reasoning)
|
||||
|
||||
The BEiT-3 **base** model can be finetuned on NLVR2 using 8 V100-32GB:
|
||||
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
|
||||
--model beit3_base_patch16_224 \
|
||||
--task nlvr2 \
|
||||
--batch_size 32 \
|
||||
--layer_decay 0.65 \
|
||||
--lr 7e-4 \
|
||||
--epochs 20 \
|
||||
--warmup_epochs 5 \
|
||||
--drop_path 0.2 \
|
||||
--sentencepiece_model /your_beit3_model_path/beit3.spm \
|
||||
--finetune /your_beit3_model_path/beit3_base_patch16_224.pth \
|
||||
--data_path /path/to/your_data \
|
||||
--output_dir /path/to/save/your_model \
|
||||
--log_dir /path/to/save/your_model/log \
|
||||
--weight_decay 0.2 \
|
||||
--seed 42 \
|
||||
--save_ckpt_freq 5 \
|
||||
--enable_deepspeed
|
||||
```
|
||||
- `--batch_size`: batch size per GPU. Effective batch size = `number of GPUs` * `--batch_size` * `--update_freq`. So in the above example, the effective batch size is `8*32 = 256`.
|
||||
- `--finetune`: weight path of your pretrained models; please download the pretrained model weights in [README.md](../README.md#pretrained-models).
|
||||
- `--enable_deepspeed`: optional. If you use apex, please enable deepspeed.
|
||||
- `--lr`: 7e-4 for `BEiT3-base`, 5e-4 for `BEiT3-base-indomain`.
|
||||
|
||||
|
||||
The BEiT-3 **large** model can be finetuned on NLVR2 using 8 V100-32GB:
|
||||
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
|
||||
--model beit3_large_patch16_224 \
|
||||
--task nlvr2 \
|
||||
--batch_size 32 \
|
||||
--layer_decay 0.85 \
|
||||
--lr 3e-4 \
|
||||
--epochs 20 \
|
||||
--warmup_epochs 5 \
|
||||
--drop_path 0.2 \
|
||||
--sentencepiece_model /your_beit3_model_path/beit3.spm \
|
||||
--finetune /your_beit3_model_path/beit3_large_patch16_224.pth \
|
||||
--data_path /path/to/your_data \
|
||||
--output_dir /path/to/save/your_model \
|
||||
--log_dir /path/to/save/your_model/log \
|
||||
--weight_decay 0.2 \
|
||||
--seed 42 \
|
||||
--save_ckpt_freq 5 \
|
||||
--enable_deepspeed \
|
||||
--checkpoint_activations
|
||||
```
|
||||
- `--batch_size`: batch size per GPU. Effective batch size = `number of GPUs` * `--batch_size` * `--update_freq`. So in the above example, the effective batch size is `8*32 = 256`.
|
||||
- `--finetune`: weight path of your pretrained models; please download the pretrained model weights in [README.md](../README.md#pretrained-models).
|
||||
- `--enable_deepspeed`: optional. If you use apex, please enable deepspeed.
|
||||
- `--lr`: 3e-4 for `BEiT3-large`, 1e-4 for `BEiT3-large-indomain`.
|
||||
- `--checkpoint_activations`: using gradient checkpointing for saving GPU memory.
|
||||
|
||||
|
||||
## Example: Evaluate BEiT-3 Finetuned model on NLVR2 (Visual Reasoning)
|
||||
|
||||
- Get the result of our fine-tuned BEiT3-base model on NLVR2 test with 8 V100-32GB:
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
|
||||
--model beit3_base_patch16_224 \
|
||||
--task nlvr2 \
|
||||
--batch_size 32 \
|
||||
--sentencepiece_model /your_beit3_model_path/beit3.spm \
|
||||
--finetune /your_beit3_model_path/beit3_base_patch16_224_nlvr2.pth \
|
||||
--data_path /path/to/your_data \
|
||||
--eval \
|
||||
--dist_eval
|
||||
```
|
||||
|
||||
Expected results:
|
||||
```
|
||||
* Acc 84.386
|
||||
```
|
||||
|
||||
- Get the result of our fine-tuned BEiT3-large model on NLVR2 test with 8 V100-32GB:
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
|
||||
--model beit3_large_patch16_224 \
|
||||
--task nlvr2 \
|
||||
--batch_size 32 \
|
||||
--sentencepiece_model /your_beit3_model_path/beit3.spm \
|
||||
--finetune /your_beit3_model_path/beit3_large_patch16_224_nlvr2.pth \
|
||||
--data_path /path/to/your_data \
|
||||
--eval \
|
||||
--dist_eval
|
||||
```
|
||||
|
||||
Expected results:
|
||||
```
|
||||
* Acc 89.437
|
||||
```
|
||||
@@ -0,0 +1,161 @@
|
||||
# Fine-tuning BEiT-3 on Image-text Retrieval
|
||||
|
||||
## COCO Retrieval Setup
|
||||
|
||||
1. [Setup environment](../README.md#setup).
|
||||
2. Download [2014 train images](http://images.cocodataset.org/zips/train2014.zip), [2014 val images](http://images.cocodataset.org/zips/val2014.zip) and [karpathy split](https://cs.stanford.edu/people/karpathy/deepimagesent/caption_datasets.zip), then organize the dataset as following structure:
|
||||
|
||||
```
|
||||
/path/to/your_data/
|
||||
train2014/
|
||||
COCO_train2014_000000000009.jpg
|
||||
...
|
||||
val2014/
|
||||
COCO_val2014_000000000042.jpg
|
||||
...
|
||||
dataset_coco.json
|
||||
```
|
||||
|
||||
We then generate the index json files using the following command. [beit3.spm](https://github.com/addf400/files/releases/download/beit3/beit3.spm) is the sentencepiece model used for tokenizing texts.
|
||||
```
|
||||
from datasets import RetrievalDataset
|
||||
from transformers import XLMRobertaTokenizer
|
||||
|
||||
tokenizer = XLMRobertaTokenizer("/your_beit3_model_path/beit3.spm")
|
||||
|
||||
RetrievalDataset.make_coco_dataset_index(
|
||||
data_path="/path/to/your_data",
|
||||
tokenizer=tokenizer,
|
||||
)
|
||||
```
|
||||
|
||||
|
||||
## Flickr30k Retrieval Setup
|
||||
|
||||
1. [Setup environment](README.md#setup).
|
||||
2. Sign [flickr images request form](https://forms.illinois.edu/sec/229675) and download [karpathy split](https://cs.stanford.edu/people/karpathy/deepimagesent/caption_datasets.zip), then organize the dataset as following structure:
|
||||
|
||||
```
|
||||
/path/to/your_data/
|
||||
flickr30k-images/
|
||||
2923475135.jpg
|
||||
...
|
||||
dataset_flickr30k.json
|
||||
```
|
||||
|
||||
We then generate the index json files using the following command. [beit3.spm](https://github.com/addf400/files/releases/download/beit3/beit3.spm) is the sentencepiece model used for tokenizing texts.
|
||||
```
|
||||
from datasets import RetrievalDataset
|
||||
from transformers import XLMRobertaTokenizer
|
||||
|
||||
tokenizer = XLMRobertaTokenizer("/your_beit3_model_path/beit3.spm")
|
||||
|
||||
RetrievalDataset.make_flickr30k_dataset_index(
|
||||
data_path="/path/to/your_data",
|
||||
tokenizer=tokenizer,
|
||||
karpathy_path="/path/to/your_data",
|
||||
)
|
||||
```
|
||||
|
||||
|
||||
## Example: Fine-tuning BEiT-3 on Retrieval
|
||||
|
||||
The BEiT-3 **base** model can be finetuned on retrieval tasks using 16 V100-32GB:
|
||||
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=16 run_beit3_finetuning.py \
|
||||
--model beit3_base_patch16_384 \
|
||||
--input_size 384 \
|
||||
--task coco_retrieval \
|
||||
--batch_size 192 \
|
||||
--layer_decay 0.65 \
|
||||
--lr 2e-4 \
|
||||
--epochs 15 \
|
||||
--warmup_epochs 3 \
|
||||
--drop_path 0.2 \
|
||||
--sentencepiece_model /your_beit3_model_path/beit3.spm \
|
||||
--finetune /your_beit3_model_path/beit3_base_itc_patch16_224.pth \
|
||||
--data_path /path/to/your_data \
|
||||
--output_dir /path/to/save/your_model \
|
||||
--log_dir /path/to/save/your_model/log \
|
||||
--weight_decay 0.05 \
|
||||
--seed 42 \
|
||||
--save_ckpt_freq 5 \
|
||||
--enable_deepspeed \
|
||||
--checkpoint_activations
|
||||
```
|
||||
- `--batch_size`: batch size per GPU. Effective batch size = `number of GPUs` * `--batch_size` * `--update_freq`. So in the above example, the effective batch size is `192*16 = 3072`.
|
||||
- `--finetune`: weight path of your pretrained models; please download the pretrained model weights in [README.md](../README.md#pretrained-models)
|
||||
- `--task`: **coco_retrieval** for COCO retrieval, **flickr30k** for Flickr30k retrieval
|
||||
- `--lr`: 2e-4 for COCO retrieval, 1e-4 for Flickr30k retrieval
|
||||
- `--epochs`: 15 for COCO retrieval, 20 for Flickr30k retrieval
|
||||
- `--warmup_epochs`: 3 for COCO retrieval, 5 for Flickr30k retrieval
|
||||
- `--checkpoint_activations`: using gradient checkpointing for saving GPU memory
|
||||
|
||||
|
||||
The BEiT-3 **large** model can be finetuned on retrieval tasks using 2x16 V100-32GB:
|
||||
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=16 --nnodes=2 --node_rank=$NODE_RANK \
|
||||
--master_addr=$MASTER_ADDR --master_port=$MASTER_PORT run_beit3_finetuning.py \
|
||||
--model beit3_large_patch16_384 \
|
||||
--input_size 384 \
|
||||
--task coco_retrieval \
|
||||
--batch_size 96 \
|
||||
--layer_decay 0.85 \
|
||||
--lr 5e-5 \
|
||||
--epochs 15 \
|
||||
--warmup_epochs 3 \
|
||||
--drop_path 0.2 \
|
||||
--sentencepiece_model /your_beit3_model_path/beit3.spm \
|
||||
--finetune /your_beit3_model_path/beit3_large_itc_patch16_224.pth \
|
||||
--data_path /path/to/your_data \
|
||||
--output_dir /path/to/save/your_model \
|
||||
--log_dir /path/to/save/your_model/log \
|
||||
--weight_decay 0.05 \
|
||||
--seed 42 \
|
||||
--save_ckpt_freq 5 \
|
||||
--enable_deepspeed \
|
||||
--checkpoint_activations
|
||||
```
|
||||
- `--batch_size`: batch size per GPU. Effective batch size = `number of GPUs` * `--batch_size` * `--update_freq`. So in the above example, the effective batch size is `96*32 = 3072`.
|
||||
- `--finetune`: weight path of your pretrained models; please download the pretrained model weights in [README.md](../README.md#pretrained-models)
|
||||
- `--task`: **coco_retrieval** for COCO retrieval, **flickr30k** for Flickr30k retrieval
|
||||
- `--epochs`: 15 for COCO retrieval, 20 for Flickr30k retrieval
|
||||
- `--warmup_epochs`: 3 for COCO retrieval, 5 for Flickr30k retrieval
|
||||
- `--checkpoint_activations`: using gradient checkpointing for saving GPU memory
|
||||
|
||||
|
||||
## Example: Evaluate BEiT-3 Fine-tuned model on COCO Retrieval and Flickr30k Retrieval
|
||||
|
||||
- Get the results of our fine-tuned BEiT3-base model on retrieval tasks using a single GPU:
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=1 run_beit3_finetuning.py \
|
||||
--model beit3_base_patch16_384 \
|
||||
--input_size 384 \
|
||||
--task coco_retrieval \
|
||||
--batch_size 16 \
|
||||
--sentencepiece_model /your_beit3_model_path/beit3.spm \
|
||||
--finetune /your_beit3_model_path/beit3_base_patch16_384_coco_retrieval.pth \
|
||||
--data_path /path/to/your_data \
|
||||
--eval \
|
||||
--dist_eval
|
||||
```
|
||||
- `--task`: **coco_retrieval** for COCO retrieval, **flickr30k** for Flickr30k retrieval
|
||||
- `--finetune`: **beit3_base_patch16_384_coco_retrieval.pth** for COCO retrieval, **beit3_base_patch16_384_f30k_retrieval.pth** for Flickr30k retrieval
|
||||
|
||||
- Get the results of our fine-tuned BEiT3-large model on retrieval tasks using a single GPU:
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=1 run_beit3_finetuning.py \
|
||||
--model beit3_large_patch16_384 \
|
||||
--input_size 384 \
|
||||
--task coco_retrieval \
|
||||
--batch_size 16 \
|
||||
--sentencepiece_model /your_beit3_model_path/beit3.spm \
|
||||
--finetune /your_beit3_model_path/beit3_large_patch16_384_coco_retrieval.pth \
|
||||
--data_path /path/to/your_data \
|
||||
--eval \
|
||||
--dist_eval
|
||||
```
|
||||
- `--task`: **coco_retrieval** for COCO retrieval, **flickr30k** for Flickr30k retrieval
|
||||
- `--finetune`: **beit3_large_patch16_384_coco_retrieval.pth** for COCO retrieval, **beit3_large_patch16_384_f30k_retrieval.pth** for Flickr30k retrieval
|
||||
@@ -0,0 +1,144 @@
|
||||
# Fine-tuning BEiT-3 on VQAv2 (Visual Question Answering)
|
||||
|
||||
|
||||
## Setup
|
||||
|
||||
1. [Setup environment](../README.md#setup).
|
||||
2. Download COCO [2014 train images](http://images.cocodataset.org/zips/train2014.zip), [2014 val images](http://images.cocodataset.org/zips/val2014.zip), [2015 test images](http://images.cocodataset.org/zips/test2015.zip), annotations ([train](https://s3.amazonaws.com/cvmlp/vqa/mscoco/vqa/v2_Annotations_Train_mscoco.zip), [val](https://s3.amazonaws.com/cvmlp/vqa/mscoco/vqa/v2_Annotations_Val_mscoco.zip)), and questions ([train](https://s3.amazonaws.com/cvmlp/vqa/mscoco/vqa/v2_Questions_Train_mscoco.zip), [val](https://s3.amazonaws.com/cvmlp/vqa/mscoco/vqa/v2_Questions_Val_mscoco.zip), [test](https://s3.amazonaws.com/cvmlp/vqa/mscoco/vqa/v2_Questions_Test_mscoco.zip)), then organize the dataset as following structure:
|
||||
|
||||
```
|
||||
/path/to/your_data/
|
||||
train2014/
|
||||
COCO_train2014_000000000009.jpg
|
||||
...
|
||||
val2014/
|
||||
COCO_val2014_000000000042.jpg
|
||||
...
|
||||
test2015/
|
||||
COCO_test2015_000000000001.jpg
|
||||
...
|
||||
vqa/
|
||||
v2_OpenEnded_mscoco_train2014_questions.json
|
||||
v2_OpenEnded_mscoco_val2014_questions.json
|
||||
v2_OpenEnded_mscoco_test2015_questions.json
|
||||
v2_OpenEnded_mscoco_test-dev2015_questions.json
|
||||
v2_mscoco_train2014_annotations.json
|
||||
v2_mscoco_val2014_annotations.json
|
||||
```
|
||||
|
||||
We then generate the index json files using the following command. [beit3.spm](https://github.com/addf400/files/releases/download/beit3/beit3.spm) is the sentencepiece model used for tokenizing texts.
|
||||
```
|
||||
from datasets import VQAv2Dataset
|
||||
from transformers import XLMRobertaTokenizer
|
||||
|
||||
tokenizer = XLMRobertaTokenizer("/your_beit3_model_path/beit3.spm")
|
||||
|
||||
VQAv2Dataset.make_dataset_index(
|
||||
data_path="/path/to/your_data",
|
||||
tokenizer=tokenizer,
|
||||
annotation_data_path="/path/to/your_data/vqa",
|
||||
)
|
||||
```
|
||||
|
||||
|
||||
## Example: Fine-tuning BEiT-3 on VQAv2 (Visual Question Answering)
|
||||
|
||||
The BEiT-3 **base** model can be finetuned on VQAv2 using 8 V100-32GB:
|
||||
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
|
||||
--model beit3_base_patch16_480 \
|
||||
--input_size 480 \
|
||||
--task vqav2 \
|
||||
--batch_size 16 \
|
||||
--layer_decay 1.0 \
|
||||
--lr 3e-5 \
|
||||
--update_freq 1 \
|
||||
--randaug \
|
||||
--epochs 10 \
|
||||
--warmup_epochs 1 \
|
||||
--drop_path 0.1 \
|
||||
--sentencepiece_model /your_beit3_model_path/beit3.spm \
|
||||
--finetune /your_beit3_model_path/beit3_base_patch16_224.pth \
|
||||
--data_path /path/to/your_data \
|
||||
--output_dir /path/to/save/your_model \
|
||||
--log_dir /path/to/save/your_model/log \
|
||||
--weight_decay 0.01 \
|
||||
--seed 42 \
|
||||
--save_ckpt_freq 5 \
|
||||
--task_head_lr_weight 20 \
|
||||
--opt_betas 0.9 0.98 \
|
||||
--enable_deepspeed
|
||||
```
|
||||
- `--batch_size`: batch size per GPU. Effective batch size = `number of GPUs` * `--batch_size` * `--update_freq`. So in the above example, the effective batch size is `8*16 = 128`.
|
||||
- `--finetune`: weight path of your pretrained models; please download the pretrained model weights in [README.md](../README.md#pretrained-models)
|
||||
- `--enable_deepspeed`: optional. If you use apex, please enable deepspeed.
|
||||
|
||||
|
||||
The BEiT-3 **large** model can be finetuned on VQAv2 using 8 V100-32GB:
|
||||
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
|
||||
--model beit3_large_patch16_480 \
|
||||
--input_size 480 \
|
||||
--task vqav2 \
|
||||
--batch_size 16 \
|
||||
--layer_decay 1.0 \
|
||||
--lr 2e-5 \
|
||||
--update_freq 1 \
|
||||
--randaug \
|
||||
--epochs 10 \
|
||||
--warmup_epochs 1 \
|
||||
--drop_path 0.15 \
|
||||
--sentencepiece_model /your_beit3_model_path/beit3.spm \
|
||||
--finetune /your_beit3_model_path/beit3_large_patch16_224.pth \
|
||||
--data_path /path/to/your_data \
|
||||
--output_dir /path/to/save/your_model \
|
||||
--log_dir /path/to/save/your_model/log \
|
||||
--weight_decay 0.01 \
|
||||
--seed 42 \
|
||||
--save_ckpt_freq 5 \
|
||||
--task_head_lr_weight 20 \
|
||||
--opt_betas 0.9 0.98 \
|
||||
--enable_deepspeed \
|
||||
--checkpoint_activations
|
||||
```
|
||||
- `--batch_size`: batch size per GPU. Effective batch size = `number of GPUs` * `--batch_size` * `--update_freq`. So in the above example, the effective batch size is `8*16 = 128`.
|
||||
- `--finetune`: weight path of your pretrained models; please download the pretrained model weights in [README.md](../README.md#pretrained-models)
|
||||
- `--enable_deepspeed`: optional. If you use apex, please enable deepspeed.
|
||||
- `--checkpoint_activations`: using gradient checkpointing for saving GPU memory
|
||||
|
||||
|
||||
## Example: Evaluate BEiT-3 Finetuned model on VQAv2 (Visual Question Answering)
|
||||
|
||||
- Get the prediction file of the fine-tuned BEiT3-base model on VQAv2 test with 8 V100-32GB:
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
|
||||
--model beit3_base_patch16_480 \
|
||||
--input_size 480 \
|
||||
--task vqav2 \
|
||||
--batch_size 16 \
|
||||
--sentencepiece_model /your_beit3_model_path/beit3.spm \
|
||||
--finetune /your_beit3_model_path/beit3_base_patch16_480_vqa.pth \
|
||||
--data_path /path/to/your_data \
|
||||
--output_dir /path/to/save/your_prediction \
|
||||
--eval \
|
||||
--dist_eval
|
||||
```
|
||||
|
||||
- Get the prediction file of the fine-tuned BEiT3-large model on VQAv2 test with 8 V100-32GB:
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=8 run_beit3_finetuning.py \
|
||||
--model beit3_large_patch16_480 \
|
||||
--input_size 480 \
|
||||
--task vqav2 \
|
||||
--batch_size 16 \
|
||||
--sentencepiece_model /your_beit3_model_path/beit3.spm \
|
||||
--finetune /your_beit3_model_path/beit3_large_patch16_480_vqa.pth \
|
||||
--data_path /path/to/your_data \
|
||||
--output_dir /path/to/save/your_prediction \
|
||||
--eval \
|
||||
--dist_eval
|
||||
```
|
||||
|
||||
Please then submit the prediction file in the `output_dir` to the [evaluation server](https://eval.ai/web/challenges/challenge-page/830/overview) to obtain the VQAv2 test-dev and test-std results.
|
||||
@@ -0,0 +1,190 @@
|
||||
import re
|
||||
|
||||
contractions = {
|
||||
"aint": "ain't",
|
||||
"arent": "aren't",
|
||||
"cant": "can't",
|
||||
"couldve": "could've",
|
||||
"couldnt": "couldn't",
|
||||
"couldn'tve": "couldn't've",
|
||||
"couldnt've": "couldn't've",
|
||||
"didnt": "didn't",
|
||||
"doesnt": "doesn't",
|
||||
"dont": "don't",
|
||||
"hadnt": "hadn't",
|
||||
"hadnt've": "hadn't've",
|
||||
"hadn'tve": "hadn't've",
|
||||
"hasnt": "hasn't",
|
||||
"havent": "haven't",
|
||||
"hed": "he'd",
|
||||
"hed've": "he'd've",
|
||||
"he'dve": "he'd've",
|
||||
"hes": "he's",
|
||||
"howd": "how'd",
|
||||
"howll": "how'll",
|
||||
"hows": "how's",
|
||||
"Id've": "I'd've",
|
||||
"I'dve": "I'd've",
|
||||
"Im": "I'm",
|
||||
"Ive": "I've",
|
||||
"isnt": "isn't",
|
||||
"itd": "it'd",
|
||||
"itd've": "it'd've",
|
||||
"it'dve": "it'd've",
|
||||
"itll": "it'll",
|
||||
"let's": "let's",
|
||||
"maam": "ma'am",
|
||||
"mightnt": "mightn't",
|
||||
"mightnt've": "mightn't've",
|
||||
"mightn'tve": "mightn't've",
|
||||
"mightve": "might've",
|
||||
"mustnt": "mustn't",
|
||||
"mustve": "must've",
|
||||
"neednt": "needn't",
|
||||
"notve": "not've",
|
||||
"oclock": "o'clock",
|
||||
"oughtnt": "oughtn't",
|
||||
"ow's'at": "'ow's'at",
|
||||
"'ows'at": "'ow's'at",
|
||||
"'ow'sat": "'ow's'at",
|
||||
"shant": "shan't",
|
||||
"shed've": "she'd've",
|
||||
"she'dve": "she'd've",
|
||||
"she's": "she's",
|
||||
"shouldve": "should've",
|
||||
"shouldnt": "shouldn't",
|
||||
"shouldnt've": "shouldn't've",
|
||||
"shouldn'tve": "shouldn't've",
|
||||
"somebody'd": "somebodyd",
|
||||
"somebodyd've": "somebody'd've",
|
||||
"somebody'dve": "somebody'd've",
|
||||
"somebodyll": "somebody'll",
|
||||
"somebodys": "somebody's",
|
||||
"someoned": "someone'd",
|
||||
"someoned've": "someone'd've",
|
||||
"someone'dve": "someone'd've",
|
||||
"someonell": "someone'll",
|
||||
"someones": "someone's",
|
||||
"somethingd": "something'd",
|
||||
"somethingd've": "something'd've",
|
||||
"something'dve": "something'd've",
|
||||
"somethingll": "something'll",
|
||||
"thats": "that's",
|
||||
"thered": "there'd",
|
||||
"thered've": "there'd've",
|
||||
"there'dve": "there'd've",
|
||||
"therere": "there're",
|
||||
"theres": "there's",
|
||||
"theyd": "they'd",
|
||||
"theyd've": "they'd've",
|
||||
"they'dve": "they'd've",
|
||||
"theyll": "they'll",
|
||||
"theyre": "they're",
|
||||
"theyve": "they've",
|
||||
"twas": "'twas",
|
||||
"wasnt": "wasn't",
|
||||
"wed've": "we'd've",
|
||||
"we'dve": "we'd've",
|
||||
"weve": "we've",
|
||||
"werent": "weren't",
|
||||
"whatll": "what'll",
|
||||
"whatre": "what're",
|
||||
"whats": "what's",
|
||||
"whatve": "what've",
|
||||
"whens": "when's",
|
||||
"whered": "where'd",
|
||||
"wheres": "where's",
|
||||
"whereve": "where've",
|
||||
"whod": "who'd",
|
||||
"whod've": "who'd've",
|
||||
"who'dve": "who'd've",
|
||||
"wholl": "who'll",
|
||||
"whos": "who's",
|
||||
"whove": "who've",
|
||||
"whyll": "why'll",
|
||||
"whyre": "why're",
|
||||
"whys": "why's",
|
||||
"wont": "won't",
|
||||
"wouldve": "would've",
|
||||
"wouldnt": "wouldn't",
|
||||
"wouldnt've": "wouldn't've",
|
||||
"wouldn'tve": "wouldn't've",
|
||||
"yall": "y'all",
|
||||
"yall'll": "y'all'll",
|
||||
"y'allll": "y'all'll",
|
||||
"yall'd've": "y'all'd've",
|
||||
"y'alld've": "y'all'd've",
|
||||
"y'all'dve": "y'all'd've",
|
||||
"youd": "you'd",
|
||||
"youd've": "you'd've",
|
||||
"you'dve": "you'd've",
|
||||
"youll": "you'll",
|
||||
"youre": "you're",
|
||||
"youve": "you've",
|
||||
}
|
||||
|
||||
manual_map = {
|
||||
"none": "0",
|
||||
"zero": "0",
|
||||
"one": "1",
|
||||
"two": "2",
|
||||
"three": "3",
|
||||
"four": "4",
|
||||
"five": "5",
|
||||
"six": "6",
|
||||
"seven": "7",
|
||||
"eight": "8",
|
||||
"nine": "9",
|
||||
"ten": "10",
|
||||
}
|
||||
articles = ["a", "an", "the"]
|
||||
period_strip = re.compile("(?!<=\d)(\.)(?!\d)")
|
||||
comma_strip = re.compile("(\d)(\,)(\d)")
|
||||
punct = [
|
||||
";",
|
||||
r"/",
|
||||
"[",
|
||||
"]",
|
||||
'"',
|
||||
"{",
|
||||
"}",
|
||||
"(",
|
||||
")",
|
||||
"=",
|
||||
"+",
|
||||
"\\",
|
||||
"_",
|
||||
"-",
|
||||
">",
|
||||
"<",
|
||||
"@",
|
||||
"`",
|
||||
",",
|
||||
"?",
|
||||
"!",
|
||||
]
|
||||
|
||||
|
||||
def normalize_word(token):
|
||||
_token = token
|
||||
for p in punct:
|
||||
if (p + " " in token or " " + p in token) or (
|
||||
re.search(comma_strip, token) != None
|
||||
):
|
||||
_token = _token.replace(p, "")
|
||||
else:
|
||||
_token = _token.replace(p, " ")
|
||||
token = period_strip.sub("", _token, re.UNICODE)
|
||||
|
||||
_token = []
|
||||
temp = token.lower().split()
|
||||
for word in temp:
|
||||
word = manual_map.setdefault(word, word)
|
||||
if word not in articles:
|
||||
_token.append(word)
|
||||
for i, word in enumerate(_token):
|
||||
if word in contractions:
|
||||
_token[i] = contractions[word]
|
||||
token = " ".join(_token)
|
||||
token = token.replace(",", "")
|
||||
return token
|
||||
@@ -0,0 +1,386 @@
|
||||
# --------------------------------------------------------
|
||||
# Image as a Foreign Language: BEiT Pretraining for Vision and Vision-Language Tasks (https://arxiv.org/abs/2208.10442)
|
||||
# Github source: https://github.com/microsoft/unilm/tree/master/beit3
|
||||
# Copyright (c) 2023 Microsoft
|
||||
# Licensed under The MIT License [see LICENSE for details]
|
||||
# --------------------------------------------------------'
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from timm.models.registry import register_model
|
||||
import numpy as np
|
||||
|
||||
import utils
|
||||
from modeling_utils import BEiT3Wrapper, _get_base_config, _get_large_config
|
||||
|
||||
|
||||
class TwoLayerMLP(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_features,
|
||||
hidden_features,
|
||||
out_features,
|
||||
norm_layer,
|
||||
norm_input=True,
|
||||
):
|
||||
super().__init__()
|
||||
self.norm1 = norm_layer(in_features) if norm_input else nn.Identity()
|
||||
self.dense1 = nn.Linear(in_features, hidden_features)
|
||||
self.norm2 = norm_layer(hidden_features)
|
||||
self.act = nn.GELU()
|
||||
self.dense2 = nn.Linear(hidden_features, out_features)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.norm1(x)
|
||||
x = self.dense1(x)
|
||||
x = self.norm2(x)
|
||||
x = self.act(x)
|
||||
return self.dense2(x)
|
||||
|
||||
|
||||
class Pooler(nn.Module):
|
||||
def __init__(self, input_features, output_features, norm_layer):
|
||||
super().__init__()
|
||||
self.norm = norm_layer(input_features)
|
||||
self.dense = nn.Linear(input_features, output_features)
|
||||
self.activation = nn.Tanh()
|
||||
|
||||
def forward(self, x):
|
||||
cls_rep = x[:, 0, :]
|
||||
cls_rep = self.norm(cls_rep)
|
||||
pooled_output = self.dense(cls_rep)
|
||||
pooled_output = self.activation(pooled_output)
|
||||
return pooled_output
|
||||
|
||||
|
||||
class BEiT3ForVisualReasoning(BEiT3Wrapper):
|
||||
def __init__(
|
||||
self,
|
||||
args,
|
||||
num_classes,
|
||||
norm_layer=nn.LayerNorm,
|
||||
**kwargs
|
||||
):
|
||||
super(BEiT3ForVisualReasoning, self).__init__(args=args)
|
||||
embed_dim = args.encoder_embed_dim
|
||||
self.head = TwoLayerMLP(
|
||||
in_features=embed_dim * 4,
|
||||
hidden_features=embed_dim * 2,
|
||||
out_features=num_classes,
|
||||
norm_layer=norm_layer,
|
||||
)
|
||||
init_scale = 0.001
|
||||
self.head.apply(self._init_weights)
|
||||
if isinstance(self.head.dense1, nn.Linear):
|
||||
self.head.dense1.weight.data.mul_(init_scale)
|
||||
self.head.dense1.bias.data.mul_(init_scale)
|
||||
|
||||
if isinstance(self.head.dense2, nn.Linear):
|
||||
self.head.dense2.weight.data.mul_(init_scale)
|
||||
self.head.dense2.bias.data.mul_(init_scale)
|
||||
|
||||
def forward(self, image_a, image_b, text_description, padding_mask, **kwargs):
|
||||
bsz, _ = text_description.size()
|
||||
|
||||
vision_input = torch.cat((image_a, image_b), dim=0)
|
||||
language_input = torch.cat((text_description, text_description), dim=0)
|
||||
padding_mask = torch.cat((padding_mask, padding_mask), dim=0)
|
||||
|
||||
outputs = self.beit3(
|
||||
textual_tokens=language_input,
|
||||
visual_tokens=vision_input,
|
||||
text_padding_position=padding_mask,
|
||||
)
|
||||
x = outputs["encoder_out"]
|
||||
multiway_split_position = outputs["multiway_split_position"]
|
||||
|
||||
vision_cls = x[:, 0, :]
|
||||
language_cls = x[:, multiway_split_position, :]
|
||||
cls_rep = torch.cat((vision_cls, language_cls), dim=-1)
|
||||
a, b = torch.split(cls_rep, split_size_or_sections=[bsz, bsz], dim=0)
|
||||
cls_rep = torch.cat((a, b), dim=-1)
|
||||
return self.head(cls_rep)
|
||||
|
||||
|
||||
class BEiT3ForImageClassification(BEiT3Wrapper):
|
||||
def __init__(
|
||||
self,
|
||||
args,
|
||||
num_classes,
|
||||
norm_layer=nn.LayerNorm,
|
||||
**kwargs
|
||||
):
|
||||
super(BEiT3ForImageClassification, self).__init__(args=args)
|
||||
embed_dim = args.encoder_embed_dim
|
||||
self.fc_norm = norm_layer(embed_dim)
|
||||
self.head = nn.Linear(embed_dim, num_classes) if num_classes > 0 else nn.Identity()
|
||||
|
||||
self.fc_norm.apply(self._init_weights)
|
||||
self.head.apply(self._init_weights)
|
||||
init_scale = 0.001
|
||||
if isinstance(self.head, nn.Linear):
|
||||
self.head.weight.data.mul_(init_scale)
|
||||
self.head.bias.data.mul_(init_scale)
|
||||
|
||||
def forward(self, image, **kwargs):
|
||||
x = self.beit3(textual_tokens=None, visual_tokens=image)["encoder_out"]
|
||||
t = x[:, 1:, :]
|
||||
cls_x = self.fc_norm(t.mean(1))
|
||||
return self.head(cls_x)
|
||||
|
||||
|
||||
class BEiT3ForCaptioning(BEiT3Wrapper):
|
||||
def __init__(
|
||||
self,
|
||||
args,
|
||||
**kwargs
|
||||
):
|
||||
super(BEiT3ForCaptioning, self).__init__(args=args)
|
||||
embed_dim = args.encoder_embed_dim
|
||||
self.mlm_head = nn.Linear(embed_dim, args.vocab_size)
|
||||
self.mlm_head.apply(self._init_weights)
|
||||
|
||||
def forward(self, image, text_ids, padding_mask, language_masked_pos, text_len=None, incremental_state=None, **kwargs):
|
||||
text_len = text_len if text_len is not None else text_ids.size(1)
|
||||
image_len = self.beit3.vision_embed.num_position_embeddings()
|
||||
max_len = text_len + image_len
|
||||
uni_mask = torch.zeros((max_len, max_len), dtype=torch.long, device=text_ids.device)
|
||||
i_start, i_end = 0, image_len
|
||||
t_start, t_end = image_len, max_len
|
||||
# triangle mask for caption to caption
|
||||
uni_mask[t_start:t_end, t_start:t_end] = torch.tril(torch.ones(text_len, text_len, dtype=torch.long, device=text_ids.device))
|
||||
# full attention for caption to image
|
||||
uni_mask[t_start:t_end, i_start:i_end] = 1
|
||||
# full attention for image to image
|
||||
uni_mask[i_start:i_end, i_start:i_end] = 1
|
||||
uni_mask = 1-uni_mask
|
||||
|
||||
if incremental_state is not None:
|
||||
for idx in range(self.get_num_layers()):
|
||||
if idx not in incremental_state:
|
||||
incremental_state[idx] = {}
|
||||
|
||||
# for incremental decoding
|
||||
positions = None
|
||||
if image is None:
|
||||
uni_mask = uni_mask[-2:]
|
||||
padding_mask = None
|
||||
# start position (2 (fairseq starts at 2) + cur_position) is equal to text_len
|
||||
positions = torch.arange(text_len, text_ids.size(1) + text_len, device=text_ids.device).long().unsqueeze(0)
|
||||
|
||||
outputs = self.beit3(
|
||||
textual_tokens=text_ids,
|
||||
visual_tokens=image,
|
||||
text_padding_position=padding_mask,
|
||||
attn_mask=uni_mask,
|
||||
incremental_state=incremental_state,
|
||||
positions=positions,
|
||||
)
|
||||
if image is not None:
|
||||
text_feats = outputs["encoder_out"][:, image_len:]
|
||||
else:
|
||||
text_feats = outputs["encoder_out"]
|
||||
|
||||
if language_masked_pos is not None:
|
||||
text_feats = text_feats[language_masked_pos.bool()]
|
||||
|
||||
return self.mlm_head(text_feats), incremental_state
|
||||
|
||||
|
||||
class BEiT3ForVisualQuestionAnswering(BEiT3Wrapper):
|
||||
def __init__(
|
||||
self,
|
||||
args,
|
||||
num_classes,
|
||||
norm_layer=nn.LayerNorm,
|
||||
**kwargs
|
||||
):
|
||||
super(BEiT3ForVisualQuestionAnswering, self).__init__(args=args)
|
||||
embed_dim = args.encoder_embed_dim
|
||||
self.pooler = Pooler(
|
||||
input_features=embed_dim,
|
||||
output_features=embed_dim,
|
||||
norm_layer=norm_layer,
|
||||
)
|
||||
self.pooler.apply(self._init_weights)
|
||||
self.head = nn.Sequential(
|
||||
nn.Linear(embed_dim, embed_dim * 2),
|
||||
norm_layer(embed_dim * 2),
|
||||
nn.GELU(),
|
||||
nn.Linear(embed_dim * 2, num_classes),
|
||||
)
|
||||
self.head.apply(self._init_weights)
|
||||
|
||||
def forward(self, image, question, padding_mask, **kwargs):
|
||||
outputs = self.beit3(
|
||||
textual_tokens=question,
|
||||
visual_tokens=image,
|
||||
text_padding_position=padding_mask,
|
||||
)
|
||||
x = outputs["encoder_out"]
|
||||
cls_rep = self.pooler(x)
|
||||
return self.head(cls_rep)
|
||||
|
||||
|
||||
class BEiT3ForRetrieval(BEiT3Wrapper):
|
||||
def __init__(
|
||||
self,
|
||||
args,
|
||||
**kwargs
|
||||
):
|
||||
super(BEiT3ForRetrieval, self).__init__(args=args)
|
||||
embed_dim = args.encoder_embed_dim
|
||||
self.language_head = nn.Linear(embed_dim, embed_dim, bias=False)
|
||||
self.vision_head = nn.Linear(embed_dim, embed_dim, bias=False)
|
||||
self.language_head.apply(self._init_weights)
|
||||
self.vision_head.apply(self._init_weights)
|
||||
self.criterion = utils.ClipLoss(
|
||||
rank=utils.get_rank(),
|
||||
world_size=utils.get_world_size(),
|
||||
)
|
||||
self.logit_scale = nn.Parameter(torch.ones([]) * np.log(1 / 0.07))
|
||||
|
||||
def forward(self, image=None, text_description=None, padding_mask=None, only_infer=False, **kwargs):
|
||||
if image is not None:
|
||||
outputs = self.beit3(
|
||||
textual_tokens=None,
|
||||
visual_tokens=image,
|
||||
text_padding_position=None,
|
||||
)
|
||||
x = outputs["encoder_out"]
|
||||
vision_cls = self.vision_head(x[:, 0, :])
|
||||
vision_cls = F.normalize(vision_cls, dim=-1)
|
||||
else:
|
||||
vision_cls = None
|
||||
|
||||
if text_description is not None:
|
||||
outputs = self.beit3(
|
||||
textual_tokens=text_description,
|
||||
visual_tokens=None,
|
||||
text_padding_position=padding_mask,
|
||||
)
|
||||
x = outputs["encoder_out"]
|
||||
language_cls = self.language_head(x[:, 0, :])
|
||||
language_cls = F.normalize(language_cls, dim=-1)
|
||||
else:
|
||||
language_cls = None
|
||||
|
||||
if only_infer:
|
||||
return vision_cls, language_cls
|
||||
else:
|
||||
loss, logits_per_image, logits_per_text = self.criterion(
|
||||
vision_cls, language_cls, self.logit_scale.exp())
|
||||
return loss, vision_cls, language_cls
|
||||
|
||||
|
||||
@register_model
|
||||
def beit3_base_patch16_224_imageclassification(pretrained=False, **kwargs):
|
||||
args = _get_base_config(**kwargs)
|
||||
args.normalize_output = False
|
||||
model = BEiT3ForImageClassification(args, num_classes=1000, **kwargs)
|
||||
return model
|
||||
|
||||
|
||||
@register_model
|
||||
def beit3_large_patch16_224_imageclassification(pretrained=False, **kwargs):
|
||||
args = _get_large_config(**kwargs)
|
||||
args.normalize_output = False
|
||||
model = BEiT3ForImageClassification(args, num_classes=1000, **kwargs)
|
||||
return model
|
||||
|
||||
|
||||
@register_model
|
||||
def beit3_base_patch16_224_nlvr2(pretrained=False, **kwargs):
|
||||
args = _get_base_config(**kwargs)
|
||||
model = BEiT3ForVisualReasoning(args, num_classes=2, **kwargs)
|
||||
return model
|
||||
|
||||
|
||||
@register_model
|
||||
def beit3_large_patch16_224_nlvr2(pretrained=False, **kwargs):
|
||||
args = _get_large_config(**kwargs)
|
||||
model = BEiT3ForVisualReasoning(args, num_classes=2, **kwargs)
|
||||
return model
|
||||
|
||||
|
||||
@register_model
|
||||
def beit3_base_patch16_384_vqav2(pretrained=False, **kwargs):
|
||||
args = _get_base_config(img_size=384, **kwargs)
|
||||
args.normalize_output = False
|
||||
model = BEiT3ForVisualQuestionAnswering(args, num_classes=3129, **kwargs)
|
||||
return model
|
||||
|
||||
|
||||
@register_model
|
||||
def beit3_base_patch16_480_vqav2(pretrained=False, **kwargs):
|
||||
args = _get_base_config(img_size=480, **kwargs)
|
||||
args.normalize_output = False
|
||||
model = BEiT3ForVisualQuestionAnswering(args, num_classes=3129, **kwargs)
|
||||
return model
|
||||
|
||||
|
||||
@register_model
|
||||
def beit3_large_patch16_384_vqav2(pretrained=False, **kwargs):
|
||||
args = _get_large_config(img_size=384, **kwargs)
|
||||
args.normalize_output = False
|
||||
model = BEiT3ForVisualQuestionAnswering(args, num_classes=3129, **kwargs)
|
||||
return model
|
||||
|
||||
|
||||
@register_model
|
||||
def beit3_large_patch16_480_vqav2(pretrained=False, **kwargs):
|
||||
args = _get_large_config(img_size=480, **kwargs)
|
||||
args.normalize_output = False
|
||||
model = BEiT3ForVisualQuestionAnswering(args, num_classes=3129, **kwargs)
|
||||
return model
|
||||
|
||||
|
||||
@register_model
|
||||
def beit3_large_patch16_768_vqav2(pretrained=False, **kwargs):
|
||||
args = _get_large_config(img_size=768, **kwargs)
|
||||
args.normalize_output = False
|
||||
model = BEiT3ForVisualQuestionAnswering(args, num_classes=3129, **kwargs)
|
||||
return model
|
||||
|
||||
|
||||
@register_model
|
||||
def beit3_base_patch16_224_captioning(pretrained=False, **kwargs):
|
||||
args = _get_base_config(**kwargs)
|
||||
model = BEiT3ForCaptioning(args, **kwargs)
|
||||
return model
|
||||
|
||||
|
||||
@register_model
|
||||
def beit3_base_patch16_480_captioning(pretrained=False, **kwargs):
|
||||
args = _get_base_config(img_size=480, **kwargs)
|
||||
model = BEiT3ForCaptioning(args, **kwargs)
|
||||
return model
|
||||
|
||||
|
||||
@register_model
|
||||
def beit3_large_patch16_480_captioning(pretrained=False, **kwargs):
|
||||
args = _get_large_config(img_size=480, **kwargs)
|
||||
model = BEiT3ForCaptioning(args, **kwargs)
|
||||
return model
|
||||
|
||||
|
||||
@register_model
|
||||
def beit3_base_patch16_224_retrieval(pretrained=False, **kwargs):
|
||||
args = _get_base_config(**kwargs)
|
||||
model = BEiT3ForRetrieval(args, **kwargs)
|
||||
return model
|
||||
|
||||
|
||||
@register_model
|
||||
def beit3_base_patch16_384_retrieval(pretrained=False, **kwargs):
|
||||
args = _get_base_config(img_size=384, **kwargs)
|
||||
model = BEiT3ForRetrieval(args, **kwargs)
|
||||
return model
|
||||
|
||||
|
||||
@register_model
|
||||
def beit3_large_patch16_384_retrieval(pretrained=False, **kwargs):
|
||||
args = _get_large_config(img_size=384, **kwargs)
|
||||
model = BEiT3ForRetrieval(args, **kwargs)
|
||||
return model
|
||||
@@ -0,0 +1,76 @@
|
||||
# --------------------------------------------------------
|
||||
# Image as a Foreign Language: BEiT Pretraining for Vision and Vision-Language Tasks (https://arxiv.org/abs/2208.10442)
|
||||
# Github source: https://github.com/microsoft/unilm/tree/master/beit3
|
||||
# Copyright (c) 2023 Microsoft
|
||||
# Licensed under The MIT License [see LICENSE for details]
|
||||
# --------------------------------------------------------'
|
||||
|
||||
import math
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from timm.models.layers import trunc_normal_ as __call_trunc_normal_
|
||||
|
||||
from torchscale.model.BEiT3 import BEiT3
|
||||
from torchscale.architecture.config import EncoderConfig
|
||||
|
||||
|
||||
def trunc_normal_(tensor, mean=0., std=1.):
|
||||
__call_trunc_normal_(tensor, mean=mean, std=std, a=-std, b=std)
|
||||
|
||||
|
||||
def _get_base_config(
|
||||
img_size=224, patch_size=16, drop_path_rate=0,
|
||||
checkpoint_activations=None, mlp_ratio=4, vocab_size=64010, **kwargs
|
||||
):
|
||||
return EncoderConfig(
|
||||
img_size=img_size, patch_size=patch_size, vocab_size=vocab_size, multiway=True,
|
||||
layernorm_embedding=False, normalize_output=True, no_output_layer=True,
|
||||
drop_path_rate=drop_path_rate, encoder_embed_dim=768, encoder_attention_heads=12,
|
||||
encoder_ffn_embed_dim=int(768 * mlp_ratio), encoder_layers=12,
|
||||
checkpoint_activations=checkpoint_activations,
|
||||
)
|
||||
|
||||
|
||||
def _get_large_config(
|
||||
img_size=224, patch_size=16, drop_path_rate=0,
|
||||
checkpoint_activations=None, mlp_ratio=4, vocab_size=64010, **kwargs
|
||||
):
|
||||
return EncoderConfig(
|
||||
img_size=img_size, patch_size=patch_size, vocab_size=vocab_size, multiway=True,
|
||||
layernorm_embedding=False, normalize_output=True, no_output_layer=True,
|
||||
drop_path_rate=drop_path_rate, encoder_embed_dim=1024, encoder_attention_heads=16,
|
||||
encoder_ffn_embed_dim=int(1024 * mlp_ratio), encoder_layers=24,
|
||||
checkpoint_activations=checkpoint_activations,
|
||||
)
|
||||
|
||||
|
||||
class BEiT3Wrapper(nn.Module):
|
||||
def __init__(self, args, **kwargs):
|
||||
super().__init__()
|
||||
self.args = args
|
||||
self.beit3 = BEiT3(args)
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def fix_init_weight(self):
|
||||
def rescale(param, layer_id):
|
||||
param.div_(math.sqrt(2.0 * layer_id))
|
||||
|
||||
for layer_id, layer in enumerate(self.blocks):
|
||||
rescale(layer.attn.proj.weight.data, layer_id + 1)
|
||||
rescale(layer.mlp.fc2.weight.data, layer_id + 1)
|
||||
|
||||
def get_num_layers(self):
|
||||
return self.beit3.encoder.num_layers
|
||||
|
||||
@torch.jit.ignore
|
||||
def no_weight_decay(self):
|
||||
return {'pos_embed', 'cls_token', 'beit3.encoder.embed_positions.A.weight', 'beit3.vision_embed.cls_token', 'logit_scale'}
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
trunc_normal_(m.weight, std=.02)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
@@ -0,0 +1,128 @@
|
||||
# --------------------------------------------------------
|
||||
# Image as a Foreign Language: BEiT Pretraining for Vision and Vision-Language Tasks (https://arxiv.org/abs/2208.10442)
|
||||
# Github source: https://github.com/microsoft/unilm/tree/master/beit3
|
||||
# Copyright (c) 2023 Microsoft
|
||||
# Licensed under The MIT License [see LICENSE for details]
|
||||
# --------------------------------------------------------'
|
||||
|
||||
from torch import optim as optim
|
||||
from timm.optim.lookahead import Lookahead
|
||||
|
||||
import json
|
||||
|
||||
|
||||
def get_num_layer_for_vit(var_name, num_max_layer):
|
||||
if "embed" in var_name:
|
||||
return 0
|
||||
elif var_name in (
|
||||
"cls_token", "mask_token", "pos_embed", "language_pos_embed",
|
||||
"word_embeddings.weight", "vision_cls_token", "vision_pos_embed"
|
||||
):
|
||||
return 0
|
||||
elif var_name.startswith("patch_embed"):
|
||||
return 0
|
||||
elif var_name.startswith("rel_pos_bias"):
|
||||
return num_max_layer - 1
|
||||
elif "layers." in var_name:
|
||||
layer_id = int(var_name.split('layers.')[1].split('.')[0])
|
||||
return layer_id + 1
|
||||
else:
|
||||
return num_max_layer - 1
|
||||
|
||||
|
||||
def get_is_head_flag_for_vit(var_name, num_max_layer):
|
||||
if var_name.startswith("head"):
|
||||
return 1
|
||||
# elif var_name.startswith("pooler"):
|
||||
# return 1
|
||||
else:
|
||||
return 0
|
||||
|
||||
|
||||
class LayerDecayValueAssigner(object):
|
||||
def __init__(self, values, scale_handler=None):
|
||||
self.scale_handler = scale_handler or get_num_layer_for_vit
|
||||
self.values = values
|
||||
|
||||
def get_scale(self, layer_id):
|
||||
return self.values[layer_id]
|
||||
|
||||
def get_layer_id(self, var_name):
|
||||
return self.scale_handler(var_name, len(self.values))
|
||||
|
||||
|
||||
# The implementation code is modified from Timm (https://github.com/huggingface/pytorch-image-models/tree/main/timm
|
||||
def get_parameter_groups(model, weight_decay=1e-5, skip_list=(), get_num_layer=None, get_layer_scale=None):
|
||||
parameter_group_names = {}
|
||||
parameter_group_vars = {}
|
||||
|
||||
for name, param in model.named_parameters():
|
||||
if not param.requires_grad:
|
||||
continue # frozen weights
|
||||
if len(param.shape) == 1 or name.endswith(".bias") or name in skip_list:
|
||||
group_name = "no_decay"
|
||||
this_weight_decay = 0.
|
||||
else:
|
||||
group_name = "decay"
|
||||
this_weight_decay = weight_decay
|
||||
if get_num_layer is not None:
|
||||
layer_id = get_num_layer(name)
|
||||
group_name = "layer_%d_%s" % (layer_id, group_name)
|
||||
else:
|
||||
layer_id = None
|
||||
|
||||
if group_name not in parameter_group_names:
|
||||
if get_layer_scale is not None:
|
||||
scale = get_layer_scale(layer_id)
|
||||
else:
|
||||
scale = 1.
|
||||
|
||||
parameter_group_names[group_name] = {
|
||||
"weight_decay": this_weight_decay,
|
||||
"params": [],
|
||||
"lr_scale": scale
|
||||
}
|
||||
parameter_group_vars[group_name] = {
|
||||
"weight_decay": this_weight_decay,
|
||||
"params": [],
|
||||
"lr_scale": scale
|
||||
}
|
||||
|
||||
parameter_group_vars[group_name]["params"].append(param)
|
||||
parameter_group_names[group_name]["params"].append(name)
|
||||
print("Param groups = %s" % json.dumps(parameter_group_names, indent=2))
|
||||
return list(parameter_group_vars.values())
|
||||
|
||||
|
||||
def create_optimizer(args, model, get_num_layer=None, get_layer_scale=None, filter_bias_and_bn=True, skip_list=None):
|
||||
opt_lower = args.opt.lower()
|
||||
weight_decay = args.weight_decay
|
||||
if weight_decay and filter_bias_and_bn:
|
||||
skip = {}
|
||||
if skip_list is not None:
|
||||
skip = skip_list
|
||||
elif hasattr(model, 'no_weight_decay'):
|
||||
skip = model.no_weight_decay()
|
||||
parameters = get_parameter_groups(model, weight_decay, skip, get_num_layer, get_layer_scale)
|
||||
weight_decay = 0.
|
||||
else:
|
||||
parameters = model.parameters()
|
||||
|
||||
opt_args = dict(lr=args.lr, weight_decay=weight_decay)
|
||||
if hasattr(args, 'opt_eps') and args.opt_eps is not None:
|
||||
opt_args['eps'] = args.opt_eps
|
||||
if hasattr(args, 'opt_betas') and args.opt_betas is not None:
|
||||
opt_args['betas'] = args.opt_betas
|
||||
|
||||
opt_split = opt_lower.split('_')
|
||||
opt_lower = opt_split[-1]
|
||||
if opt_lower == 'adamw':
|
||||
optimizer = optim.AdamW(parameters, **opt_args)
|
||||
else:
|
||||
raise ValueError("Invalid optimizer")
|
||||
|
||||
if len(opt_split) > 1:
|
||||
if opt_split[0] == 'lookahead':
|
||||
optimizer = Lookahead(optimizer)
|
||||
|
||||
return optimizer
|
||||
@@ -0,0 +1,340 @@
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
|
||||
## aug functions
|
||||
def identity_func(img):
|
||||
return img
|
||||
|
||||
|
||||
def autocontrast_func(img, cutoff=0):
|
||||
'''
|
||||
same output as PIL.ImageOps.autocontrast
|
||||
'''
|
||||
n_bins = 256
|
||||
|
||||
def tune_channel(ch):
|
||||
n = ch.size
|
||||
cut = cutoff * n // 100
|
||||
if cut == 0:
|
||||
high, low = ch.max(), ch.min()
|
||||
else:
|
||||
hist = cv2.calcHist([ch], [0], None, [n_bins], [0, n_bins])
|
||||
low = np.argwhere(np.cumsum(hist) > cut)
|
||||
low = 0 if low.shape[0] == 0 else low[0]
|
||||
high = np.argwhere(np.cumsum(hist[::-1]) > cut)
|
||||
high = n_bins - 1 if high.shape[0] == 0 else n_bins - 1 - high[0]
|
||||
if high <= low:
|
||||
table = np.arange(n_bins)
|
||||
else:
|
||||
scale = (n_bins - 1) / (high - low)
|
||||
offset = -low * scale
|
||||
table = np.arange(n_bins) * scale + offset
|
||||
table[table < 0] = 0
|
||||
table[table > n_bins - 1] = n_bins - 1
|
||||
table = table.clip(0, 255).astype(np.uint8)
|
||||
return table[ch]
|
||||
|
||||
channels = [tune_channel(ch) for ch in cv2.split(img)]
|
||||
out = cv2.merge(channels)
|
||||
return out
|
||||
|
||||
|
||||
def equalize_func(img):
|
||||
'''
|
||||
same output as PIL.ImageOps.equalize
|
||||
PIL's implementation is different from cv2.equalize
|
||||
'''
|
||||
n_bins = 256
|
||||
|
||||
def tune_channel(ch):
|
||||
hist = cv2.calcHist([ch], [0], None, [n_bins], [0, n_bins])
|
||||
non_zero_hist = hist[hist != 0].reshape(-1)
|
||||
step = np.sum(non_zero_hist[:-1]) // (n_bins - 1)
|
||||
if step == 0: return ch
|
||||
n = np.empty_like(hist)
|
||||
n[0] = step // 2
|
||||
n[1:] = hist[:-1]
|
||||
table = (np.cumsum(n) // step).clip(0, 255).astype(np.uint8)
|
||||
return table[ch]
|
||||
|
||||
channels = [tune_channel(ch) for ch in cv2.split(img)]
|
||||
out = cv2.merge(channels)
|
||||
return out
|
||||
|
||||
|
||||
def rotate_func(img, degree, fill=(0, 0, 0)):
|
||||
'''
|
||||
like PIL, rotate by degree, not radians
|
||||
'''
|
||||
H, W = img.shape[0], img.shape[1]
|
||||
center = W / 2, H / 2
|
||||
M = cv2.getRotationMatrix2D(center, degree, 1)
|
||||
out = cv2.warpAffine(img, M, (W, H), borderValue=fill)
|
||||
return out
|
||||
|
||||
|
||||
def solarize_func(img, thresh=128):
|
||||
'''
|
||||
same output as PIL.ImageOps.posterize
|
||||
'''
|
||||
table = np.array([el if el < thresh else 255 - el for el in range(256)])
|
||||
table = table.clip(0, 255).astype(np.uint8)
|
||||
out = table[img]
|
||||
return out
|
||||
|
||||
|
||||
def color_func(img, factor):
|
||||
'''
|
||||
same output as PIL.ImageEnhance.Color
|
||||
'''
|
||||
## implementation according to PIL definition, quite slow
|
||||
# degenerate = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)[:, :, np.newaxis]
|
||||
# out = blend(degenerate, img, factor)
|
||||
# M = (
|
||||
# np.eye(3) * factor
|
||||
# + np.float32([0.114, 0.587, 0.299]).reshape(3, 1) * (1. - factor)
|
||||
# )[np.newaxis, np.newaxis, :]
|
||||
M = (
|
||||
np.float32([
|
||||
[0.886, -0.114, -0.114],
|
||||
[-0.587, 0.413, -0.587],
|
||||
[-0.299, -0.299, 0.701]]) * factor
|
||||
+ np.float32([[0.114], [0.587], [0.299]])
|
||||
)
|
||||
out = np.matmul(img, M).clip(0, 255).astype(np.uint8)
|
||||
return out
|
||||
|
||||
|
||||
def contrast_func(img, factor):
|
||||
"""
|
||||
same output as PIL.ImageEnhance.Contrast
|
||||
"""
|
||||
mean = np.sum(np.mean(img, axis=(0, 1)) * np.array([0.114, 0.587, 0.299]))
|
||||
table = np.array([(
|
||||
el - mean) * factor + mean
|
||||
for el in range(256)
|
||||
]).clip(0, 255).astype(np.uint8)
|
||||
out = table[img]
|
||||
return out
|
||||
|
||||
|
||||
def brightness_func(img, factor):
|
||||
'''
|
||||
same output as PIL.ImageEnhance.Contrast
|
||||
'''
|
||||
table = (np.arange(256, dtype=np.float32) * factor).clip(0, 255).astype(np.uint8)
|
||||
out = table[img]
|
||||
return out
|
||||
|
||||
|
||||
def sharpness_func(img, factor):
|
||||
'''
|
||||
The differences the this result and PIL are all on the 4 boundaries, the center
|
||||
areas are same
|
||||
'''
|
||||
kernel = np.ones((3, 3), dtype=np.float32)
|
||||
kernel[1][1] = 5
|
||||
kernel /= 13
|
||||
degenerate = cv2.filter2D(img, -1, kernel)
|
||||
if factor == 0.0:
|
||||
out = degenerate
|
||||
elif factor == 1.0:
|
||||
out = img
|
||||
else:
|
||||
out = img.astype(np.float32)
|
||||
degenerate = degenerate.astype(np.float32)[1:-1, 1:-1, :]
|
||||
out[1:-1, 1:-1, :] = degenerate + factor * (out[1:-1, 1:-1, :] - degenerate)
|
||||
out = out.astype(np.uint8)
|
||||
return out
|
||||
|
||||
|
||||
def shear_x_func(img, factor, fill=(0, 0, 0)):
|
||||
H, W = img.shape[0], img.shape[1]
|
||||
M = np.float32([[1, factor, 0], [0, 1, 0]])
|
||||
out = cv2.warpAffine(img, M, (W, H), borderValue=fill, flags=cv2.INTER_LINEAR).astype(np.uint8)
|
||||
return out
|
||||
|
||||
|
||||
def translate_x_func(img, offset, fill=(0, 0, 0)):
|
||||
'''
|
||||
same output as PIL.Image.transform
|
||||
'''
|
||||
H, W = img.shape[0], img.shape[1]
|
||||
M = np.float32([[1, 0, -offset], [0, 1, 0]])
|
||||
out = cv2.warpAffine(img, M, (W, H), borderValue=fill, flags=cv2.INTER_LINEAR).astype(np.uint8)
|
||||
return out
|
||||
|
||||
|
||||
def translate_y_func(img, offset, fill=(0, 0, 0)):
|
||||
'''
|
||||
same output as PIL.Image.transform
|
||||
'''
|
||||
H, W = img.shape[0], img.shape[1]
|
||||
M = np.float32([[1, 0, 0], [0, 1, -offset]])
|
||||
out = cv2.warpAffine(img, M, (W, H), borderValue=fill, flags=cv2.INTER_LINEAR).astype(np.uint8)
|
||||
return out
|
||||
|
||||
|
||||
def posterize_func(img, bits):
|
||||
'''
|
||||
same output as PIL.ImageOps.posterize
|
||||
'''
|
||||
out = np.bitwise_and(img, np.uint8(255 << (8 - bits)))
|
||||
return out
|
||||
|
||||
|
||||
def shear_y_func(img, factor, fill=(0, 0, 0)):
|
||||
H, W = img.shape[0], img.shape[1]
|
||||
M = np.float32([[1, 0, 0], [factor, 1, 0]])
|
||||
out = cv2.warpAffine(img, M, (W, H), borderValue=fill, flags=cv2.INTER_LINEAR).astype(np.uint8)
|
||||
return out
|
||||
|
||||
|
||||
def cutout_func(img, pad_size, replace=(0, 0, 0)):
|
||||
replace = np.array(replace, dtype=np.uint8)
|
||||
H, W = img.shape[0], img.shape[1]
|
||||
rh, rw = np.random.random(2)
|
||||
pad_size = pad_size // 2
|
||||
ch, cw = int(rh * H), int(rw * W)
|
||||
x1, x2 = max(ch - pad_size, 0), min(ch + pad_size, H)
|
||||
y1, y2 = max(cw - pad_size, 0), min(cw + pad_size, W)
|
||||
out = img.copy()
|
||||
out[x1:x2, y1:y2, :] = replace
|
||||
return out
|
||||
|
||||
|
||||
### level to args
|
||||
def enhance_level_to_args(MAX_LEVEL):
|
||||
def level_to_args(level):
|
||||
return ((level / MAX_LEVEL) * 1.8 + 0.1,)
|
||||
return level_to_args
|
||||
|
||||
|
||||
def shear_level_to_args(MAX_LEVEL, replace_value):
|
||||
def level_to_args(level):
|
||||
level = (level / MAX_LEVEL) * 0.3
|
||||
if np.random.random() > 0.5: level = -level
|
||||
return (level, replace_value)
|
||||
|
||||
return level_to_args
|
||||
|
||||
|
||||
def translate_level_to_args(translate_const, MAX_LEVEL, replace_value):
|
||||
def level_to_args(level):
|
||||
level = (level / MAX_LEVEL) * float(translate_const)
|
||||
if np.random.random() > 0.5: level = -level
|
||||
return (level, replace_value)
|
||||
|
||||
return level_to_args
|
||||
|
||||
|
||||
def cutout_level_to_args(cutout_const, MAX_LEVEL, replace_value):
|
||||
def level_to_args(level):
|
||||
level = int((level / MAX_LEVEL) * cutout_const)
|
||||
return (level, replace_value)
|
||||
|
||||
return level_to_args
|
||||
|
||||
|
||||
def solarize_level_to_args(MAX_LEVEL):
|
||||
def level_to_args(level):
|
||||
level = int((level / MAX_LEVEL) * 256)
|
||||
return (level, )
|
||||
return level_to_args
|
||||
|
||||
|
||||
def none_level_to_args(level):
|
||||
return ()
|
||||
|
||||
|
||||
def posterize_level_to_args(MAX_LEVEL):
|
||||
def level_to_args(level):
|
||||
level = int((level / MAX_LEVEL) * 4)
|
||||
return (level, )
|
||||
return level_to_args
|
||||
|
||||
|
||||
def rotate_level_to_args(MAX_LEVEL, replace_value):
|
||||
def level_to_args(level):
|
||||
level = (level / MAX_LEVEL) * 30
|
||||
if np.random.random() < 0.5:
|
||||
level = -level
|
||||
return (level, replace_value)
|
||||
|
||||
return level_to_args
|
||||
|
||||
|
||||
func_dict = {
|
||||
'Identity': identity_func,
|
||||
'AutoContrast': autocontrast_func,
|
||||
'Equalize': equalize_func,
|
||||
'Rotate': rotate_func,
|
||||
'Solarize': solarize_func,
|
||||
'Color': color_func,
|
||||
'Contrast': contrast_func,
|
||||
'Brightness': brightness_func,
|
||||
'Sharpness': sharpness_func,
|
||||
'ShearX': shear_x_func,
|
||||
'TranslateX': translate_x_func,
|
||||
'TranslateY': translate_y_func,
|
||||
'Posterize': posterize_func,
|
||||
'ShearY': shear_y_func,
|
||||
}
|
||||
|
||||
translate_const = 10
|
||||
MAX_LEVEL = 10
|
||||
replace_value = (128, 128, 128)
|
||||
arg_dict = {
|
||||
'Identity': none_level_to_args,
|
||||
'AutoContrast': none_level_to_args,
|
||||
'Equalize': none_level_to_args,
|
||||
'Rotate': rotate_level_to_args(MAX_LEVEL, replace_value),
|
||||
'Solarize': solarize_level_to_args(MAX_LEVEL),
|
||||
'Color': enhance_level_to_args(MAX_LEVEL),
|
||||
'Contrast': enhance_level_to_args(MAX_LEVEL),
|
||||
'Brightness': enhance_level_to_args(MAX_LEVEL),
|
||||
'Sharpness': enhance_level_to_args(MAX_LEVEL),
|
||||
'ShearX': shear_level_to_args(MAX_LEVEL, replace_value),
|
||||
'TranslateX': translate_level_to_args(
|
||||
translate_const, MAX_LEVEL, replace_value
|
||||
),
|
||||
'TranslateY': translate_level_to_args(
|
||||
translate_const, MAX_LEVEL, replace_value
|
||||
),
|
||||
'Posterize': posterize_level_to_args(MAX_LEVEL),
|
||||
'ShearY': shear_level_to_args(MAX_LEVEL, replace_value),
|
||||
}
|
||||
|
||||
|
||||
class RandomAugment(object):
|
||||
|
||||
def __init__(self, N=2, M=10, isPIL=False, augs=[]):
|
||||
self.N = N
|
||||
self.M = M
|
||||
self.isPIL = isPIL
|
||||
if augs:
|
||||
self.augs = augs
|
||||
else:
|
||||
self.augs = list(arg_dict.keys())
|
||||
|
||||
def get_random_ops(self):
|
||||
sampled_ops = np.random.choice(self.augs, self.N)
|
||||
return [(op, 0.5, self.M) for op in sampled_ops]
|
||||
|
||||
def __call__(self, img):
|
||||
if self.isPIL:
|
||||
img = np.array(img)
|
||||
ops = self.get_random_ops()
|
||||
for name, prob, level in ops:
|
||||
if np.random.random() > prob:
|
||||
continue
|
||||
args = arg_dict[name](level)
|
||||
img = func_dict[name](img, *args)
|
||||
return img
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
a = RandomAugment()
|
||||
img = np.random.randn(32, 32, 3)
|
||||
a(img)
|
||||
@@ -0,0 +1,22 @@
|
||||
torch
|
||||
torchvision
|
||||
timm==0.4.12
|
||||
Pillow
|
||||
blobfile
|
||||
mypy
|
||||
numpy
|
||||
pytest
|
||||
requests
|
||||
einops
|
||||
tensorboardX
|
||||
scipy
|
||||
ftfy
|
||||
opencv-python
|
||||
sentencepiece
|
||||
pyarrow
|
||||
torchmetrics==0.7.3
|
||||
transformers
|
||||
deepspeed==0.4.0
|
||||
pycocotools
|
||||
pycocoevalcap
|
||||
torchscale==0.2.0
|
||||
@@ -0,0 +1,448 @@
|
||||
# --------------------------------------------------------
|
||||
# Image as a Foreign Language: BEiT Pretraining for Vision and Vision-Language Tasks (https://arxiv.org/abs/2208.10442)
|
||||
# Github source: https://github.com/microsoft/unilm/tree/master/beit3
|
||||
# Copyright (c) 2023 Microsoft
|
||||
# Licensed under The MIT License [see LICENSE for details]
|
||||
# --------------------------------------------------------'
|
||||
|
||||
import argparse
|
||||
import datetime
|
||||
import numpy as np
|
||||
import time
|
||||
import torch
|
||||
import torch.backends.cudnn as cudnn
|
||||
import json
|
||||
import os
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from timm.data.mixup import Mixup
|
||||
from timm.models import create_model
|
||||
from timm.utils import ModelEma
|
||||
from optim_factory import create_optimizer, get_parameter_groups, \
|
||||
LayerDecayValueAssigner, get_is_head_flag_for_vit
|
||||
|
||||
from engine_for_finetuning import train_one_epoch, get_handler, evaluate
|
||||
from datasets import create_downstream_dataset
|
||||
from utils import NativeScalerWithGradNormCount as NativeScaler
|
||||
import utils
|
||||
import modeling_finetune
|
||||
|
||||
|
||||
def get_args():
|
||||
parser = argparse.ArgumentParser('BEiT fine-tuning and evaluation script for image classification', add_help=False)
|
||||
|
||||
# Model parameters
|
||||
parser.add_argument('--model', default='beit_base_patch16_224', type=str, metavar='MODEL',
|
||||
help='Name of model to train')
|
||||
parser.add_argument('--task', type=str, required=True,
|
||||
choices=['nlvr2', 'vqav2', 'flickr30k', 'coco_retrieval', 'coco_captioning', 'nocaps', 'imagenet'],
|
||||
help='Name of task to fine-tuning')
|
||||
|
||||
parser.add_argument('--input_size', default=224, type=int,
|
||||
help='images input size')
|
||||
parser.add_argument('--drop_path', type=float, default=0.1, metavar='PCT',
|
||||
help='Drop path rate (default: 0.1)')
|
||||
|
||||
parser.add_argument('--checkpoint_activations', action='store_true', default=None,
|
||||
help='Enable checkpointing to save your memory.')
|
||||
parser.add_argument('--sentencepiece_model', type=str, required=True,
|
||||
help='Sentencepiece model path for the pretrained model.')
|
||||
parser.add_argument('--vocab_size', type=int, default=64010)
|
||||
parser.add_argument('--num_max_bpe_tokens', type=int, default=64)
|
||||
|
||||
parser.add_argument('--model_ema', action='store_true', default=False)
|
||||
parser.add_argument('--model_ema_decay', type=float, default=0.9999, help='')
|
||||
parser.add_argument('--model_ema_force_cpu', action='store_true', default=False, help='')
|
||||
|
||||
# Optimizer parameters
|
||||
parser.add_argument('--opt', default='adamw', type=str, metavar='OPTIMIZER',
|
||||
help='Optimizer (default: "adamw"')
|
||||
parser.add_argument('--opt_eps', default=1e-8, type=float, metavar='EPSILON',
|
||||
help='Optimizer Epsilon (default: 1e-8)')
|
||||
parser.add_argument('--opt_betas', default=[0.9, 0.999], type=float, nargs='+', metavar='BETA',
|
||||
help='Optimizer Betas (default: 0.9, 0.999, use opt default)')
|
||||
parser.add_argument('--clip_grad', type=float, default=None, metavar='NORM',
|
||||
help='Clip gradient norm (default: None, no clipping)')
|
||||
parser.add_argument('--momentum', type=float, default=0.9, metavar='M',
|
||||
help='SGD momentum (default: 0.9)')
|
||||
parser.add_argument('--weight_decay', type=float, default=0.05,
|
||||
help='weight decay (default: 0.05)')
|
||||
|
||||
parser.add_argument('--lr', type=float, default=5e-4, metavar='LR',
|
||||
help='learning rate (default: 5e-4)')
|
||||
parser.add_argument('--layer_decay', type=float, default=0.9)
|
||||
parser.add_argument('--task_head_lr_weight', type=float, default=0)
|
||||
|
||||
parser.add_argument('--warmup_lr', type=float, default=1e-6, metavar='LR',
|
||||
help='warmup learning rate (default: 1e-6)')
|
||||
parser.add_argument('--min_lr', type=float, default=1e-6, metavar='LR',
|
||||
help='lower lr bound for cyclic schedulers that hit 0 (1e-6)')
|
||||
parser.add_argument('--warmup_epochs', type=int, default=5, metavar='N',
|
||||
help='epochs to warmup LR, if scheduler supports')
|
||||
parser.add_argument('--warmup_steps', type=int, default=-1, metavar='N',
|
||||
help='num of steps to warmup LR, will overload warmup_epochs if set > 0')
|
||||
|
||||
parser.add_argument('--batch_size', default=64, type=int)
|
||||
parser.add_argument('--eval_batch_size', default=None, type=int)
|
||||
parser.add_argument('--epochs', default=20, type=int)
|
||||
parser.add_argument('--update_freq', default=1, type=int)
|
||||
parser.add_argument('--save_ckpt_freq', default=5, type=int)
|
||||
|
||||
# Augmentation parameters
|
||||
parser.add_argument('--randaug', action='store_true', default=False)
|
||||
parser.add_argument('--train_interpolation', type=str, default='bicubic',
|
||||
help='Training interpolation (random, bilinear, bicubic default: "bicubic")')
|
||||
|
||||
# Finetuning params
|
||||
parser.add_argument('--finetune', default='',
|
||||
help='finetune from checkpoint')
|
||||
parser.add_argument('--model_key', default='model|module', type=str)
|
||||
parser.add_argument('--model_prefix', default='', type=str)
|
||||
|
||||
# Dataset parameters
|
||||
parser.add_argument('--data_path', default='/datasets01/imagenet_full_size/061417/', type=str,
|
||||
help='dataset path')
|
||||
|
||||
parser.add_argument('--output_dir', default='',
|
||||
help='path where to save, empty for no saving')
|
||||
parser.add_argument('--log_dir', default=None,
|
||||
help='path where to tensorboard log')
|
||||
parser.add_argument('--device', default='cuda',
|
||||
help='device to use for training / testing')
|
||||
parser.add_argument('--seed', default=0, type=int)
|
||||
parser.add_argument('--resume', default='',
|
||||
help='resume from checkpoint')
|
||||
parser.add_argument('--auto_resume', action='store_true')
|
||||
parser.add_argument('--no_auto_resume', action='store_false', dest='auto_resume')
|
||||
parser.set_defaults(auto_resume=True)
|
||||
|
||||
parser.add_argument('--save_ckpt', action='store_true')
|
||||
parser.add_argument('--no_save_ckpt', action='store_false', dest='save_ckpt')
|
||||
parser.set_defaults(save_ckpt=True)
|
||||
|
||||
parser.add_argument('--start_epoch', default=0, type=int, metavar='N',
|
||||
help='start epoch')
|
||||
parser.add_argument('--eval', action='store_true',
|
||||
help='Perform evaluation only')
|
||||
parser.add_argument('--dist_eval', action='store_true', default=False,
|
||||
help='Enabling distributed evaluation')
|
||||
parser.add_argument('--num_workers', default=10, type=int)
|
||||
parser.add_argument('--pin_mem', action='store_true',
|
||||
help='Pin CPU memory in DataLoader for more efficient (sometimes) transfer to GPU.')
|
||||
parser.add_argument('--no_pin_mem', action='store_false', dest='pin_mem')
|
||||
parser.set_defaults(pin_mem=True)
|
||||
|
||||
# distributed training parameters
|
||||
parser.add_argument('--world_size', default=1, type=int,
|
||||
help='number of distributed processes')
|
||||
parser.add_argument('--local_rank', default=-1, type=int)
|
||||
parser.add_argument('--dist_on_itp', action='store_true')
|
||||
parser.add_argument('--dist_url', default='env://',
|
||||
help='url used to set up distributed training')
|
||||
|
||||
# parameter for dump predictions (VQA, COCO captioning, NoCaps)
|
||||
parser.add_argument('--task_cache_path', default=None, type=str)
|
||||
|
||||
# parameter for imagenet finetuning
|
||||
parser.add_argument('--nb_classes', default=1000, type=int,
|
||||
help='number of the classification types')
|
||||
parser.add_argument('--mixup', type=float, default=0,
|
||||
help='mixup alpha, mixup enabled if > 0.')
|
||||
parser.add_argument('--cutmix', type=float, default=0,
|
||||
help='cutmix alpha, cutmix enabled if > 0.')
|
||||
parser.add_argument('--cutmix_minmax', type=float, nargs='+', default=None,
|
||||
help='cutmix min/max ratio, overrides alpha and enables cutmix if set (default: None)')
|
||||
parser.add_argument('--mixup_prob', type=float, default=1.0,
|
||||
help='Probability of performing mixup or cutmix when either/both is enabled')
|
||||
parser.add_argument('--mixup_switch_prob', type=float, default=0.5,
|
||||
help='Probability of switching to cutmix when both mixup and cutmix enabled')
|
||||
parser.add_argument('--mixup_mode', type=str, default='batch',
|
||||
help='How to apply mixup/cutmix params. Per "batch", "pair", or "elem"')
|
||||
|
||||
# augmentation parameters for imagenet finetuning
|
||||
parser.add_argument('--color_jitter', type=float, default=0.4, metavar='PCT',
|
||||
help='Color jitter factor (default: 0.4)')
|
||||
parser.add_argument('--aa', type=str, default='rand-m9-mstd0.5-inc1', metavar='NAME',
|
||||
help='Use AutoAugment policy. "v0" or "original". " + "(default: rand-m9-mstd0.5-inc1)')
|
||||
parser.add_argument('--smoothing', type=float, default=0.1,
|
||||
help='Label smoothing (default: 0.1)')
|
||||
|
||||
# evaluation parameters for imagenet
|
||||
parser.add_argument('--crop_pct', type=float, default=None)
|
||||
|
||||
# random Erase params for imagenet finetuning
|
||||
parser.add_argument('--reprob', type=float, default=0.25, metavar='PCT',
|
||||
help='Random erase prob (default: 0.25)')
|
||||
parser.add_argument('--remode', type=str, default='pixel',
|
||||
help='Random erase mode (default: "pixel")')
|
||||
parser.add_argument('--recount', type=int, default=1,
|
||||
help='Random erase count (default: 1)')
|
||||
parser.add_argument('--resplit', action='store_true', default=False,
|
||||
help='Do not random erase first (clean) augmentation split')
|
||||
|
||||
# parameter for captioning finetuning
|
||||
parser.add_argument('--captioning_mask_prob', type=float, default=0.6)
|
||||
parser.add_argument('--drop_worst_ratio', type=float, default=0.2)
|
||||
parser.add_argument('--drop_worst_after', type=int, default=12000)
|
||||
parser.add_argument('--num_beams', type=int, default=3)
|
||||
parser.add_argument('--length_penalty', type=float, default=0.6)
|
||||
|
||||
# label smoothing for imagenet and captioning
|
||||
parser.add_argument('--label_smoothing', type=float, default=0.1)
|
||||
|
||||
# deepspeed parameters
|
||||
parser.add_argument('--enable_deepspeed', action='store_true', default=False)
|
||||
parser.add_argument('--initial_scale_power', type=int, default=16)
|
||||
parser.add_argument('--zero_stage', default=0, type=int,
|
||||
help='ZeRO optimizer stage (default: 0)')
|
||||
|
||||
known_args, _ = parser.parse_known_args()
|
||||
|
||||
if known_args.enable_deepspeed:
|
||||
try:
|
||||
import deepspeed
|
||||
from deepspeed import DeepSpeedConfig
|
||||
parser = deepspeed.add_config_arguments(parser)
|
||||
ds_init = deepspeed.initialize
|
||||
except:
|
||||
print("Please 'pip install deepspeed==0.4.0'")
|
||||
exit(0)
|
||||
else:
|
||||
ds_init = None
|
||||
|
||||
return parser.parse_args(), ds_init
|
||||
|
||||
|
||||
def main(args, ds_init):
|
||||
utils.init_distributed_mode(args)
|
||||
|
||||
if ds_init is not None:
|
||||
utils.create_ds_config(args)
|
||||
|
||||
if args.task_cache_path is None:
|
||||
args.task_cache_path = args.output_dir
|
||||
|
||||
print(args)
|
||||
|
||||
device = torch.device(args.device)
|
||||
|
||||
# fix the seed for reproducibility
|
||||
seed = args.seed + utils.get_rank()
|
||||
torch.manual_seed(seed)
|
||||
np.random.seed(seed)
|
||||
# random.seed(seed)
|
||||
|
||||
cudnn.benchmark = True
|
||||
|
||||
if utils.get_rank() == 0 and args.log_dir is not None:
|
||||
os.makedirs(args.log_dir, exist_ok=True)
|
||||
log_writer = utils.TensorboardLogger(log_dir=args.log_dir)
|
||||
else:
|
||||
log_writer = None
|
||||
|
||||
data_loader_train, data_loader_val = create_downstream_dataset(args)
|
||||
|
||||
if not args.model.endswith(args.task):
|
||||
if args.task in ("flickr30k", "coco_retrieval"):
|
||||
model_config = "%s_retrieval" % args.model
|
||||
elif args.task in ("coco_captioning", "nocaps"):
|
||||
model_config = "%s_captioning" % args.model
|
||||
elif args.task in ("imagenet"):
|
||||
model_config = "%s_imageclassification" % args.model
|
||||
else:
|
||||
model_config = "%s_%s" % (args.model, args.task)
|
||||
else:
|
||||
model_config = args.model
|
||||
print("model_config = %s" % model_config)
|
||||
model = create_model(
|
||||
model_config,
|
||||
pretrained=False,
|
||||
drop_path_rate=args.drop_path,
|
||||
vocab_size=args.vocab_size,
|
||||
checkpoint_activations=args.checkpoint_activations,
|
||||
)
|
||||
|
||||
if args.finetune:
|
||||
utils.load_model_and_may_interpolate(args.finetune, model, args.model_key, args.model_prefix)
|
||||
|
||||
model.to(device)
|
||||
|
||||
model_ema = None
|
||||
if args.model_ema:
|
||||
# Important to create EMA model after cuda(), DP wrapper, and AMP but before SyncBN and DDP wrapper
|
||||
model_ema = ModelEma(
|
||||
model,
|
||||
decay=args.model_ema_decay,
|
||||
device='cpu' if args.model_ema_force_cpu else '',
|
||||
resume='')
|
||||
print("Using EMA with decay = %.8f" % args.model_ema_decay)
|
||||
|
||||
model_without_ddp = model
|
||||
n_parameters = sum(p.numel() for p in model.parameters() if p.requires_grad)
|
||||
|
||||
print("Model = %s" % str(model_without_ddp))
|
||||
print('number of params:', n_parameters)
|
||||
|
||||
total_batch_size = args.batch_size * args.update_freq * utils.get_world_size()
|
||||
num_training_steps_per_epoch = len(data_loader_train.dataset) // total_batch_size
|
||||
print("LR = %.8f" % args.lr)
|
||||
print("Batch size = %d" % total_batch_size)
|
||||
print("Update frequent = %d" % args.update_freq)
|
||||
print("Number of training examples = %d" % len(data_loader_train.dataset))
|
||||
print("Number of training training per epoch = %d" % num_training_steps_per_epoch)
|
||||
|
||||
num_layers = model_without_ddp.get_num_layers()
|
||||
if args.layer_decay < 1.0:
|
||||
lrs = list(args.layer_decay ** (num_layers + 1 - i) for i in range(num_layers + 2))
|
||||
assigner = LayerDecayValueAssigner(lrs)
|
||||
elif args.task_head_lr_weight > 1:
|
||||
assigner = LayerDecayValueAssigner([1.0, args.task_head_lr_weight], scale_handler=get_is_head_flag_for_vit)
|
||||
else:
|
||||
assigner = None
|
||||
|
||||
if assigner is not None:
|
||||
print("Assigned values = %s" % str(assigner.values))
|
||||
|
||||
skip_weight_decay_list = model.no_weight_decay()
|
||||
|
||||
if args.distributed:
|
||||
torch.distributed.barrier()
|
||||
if args.enable_deepspeed:
|
||||
loss_scaler = None
|
||||
optimizer_params = get_parameter_groups(
|
||||
model, args.weight_decay, skip_weight_decay_list,
|
||||
assigner.get_layer_id if assigner is not None else None,
|
||||
assigner.get_scale if assigner is not None else None)
|
||||
model, optimizer, _, _ = ds_init(
|
||||
args=args, model=model, model_parameters=optimizer_params,
|
||||
dist_init_required=not args.distributed,
|
||||
)
|
||||
|
||||
print("model.gradient_accumulation_steps() = %d" % model.gradient_accumulation_steps())
|
||||
assert model.gradient_accumulation_steps() == args.update_freq
|
||||
else:
|
||||
if args.distributed:
|
||||
model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[args.gpu], find_unused_parameters=True)
|
||||
model_without_ddp = model.module
|
||||
|
||||
optimizer = create_optimizer(
|
||||
args, model_without_ddp, skip_list=skip_weight_decay_list,
|
||||
get_num_layer=assigner.get_layer_id if assigner is not None else None,
|
||||
get_layer_scale=assigner.get_scale if assigner is not None else None)
|
||||
loss_scaler = NativeScaler()
|
||||
|
||||
lr_schedule_values = utils.cosine_scheduler(
|
||||
args.lr, args.min_lr, args.epochs, num_training_steps_per_epoch,
|
||||
warmup_epochs=args.warmup_epochs, warmup_steps=args.warmup_steps,
|
||||
)
|
||||
|
||||
utils.auto_load_model(
|
||||
args=args, model=model, model_without_ddp=model_without_ddp,
|
||||
optimizer=optimizer, loss_scaler=loss_scaler, model_ema=model_ema)
|
||||
|
||||
task_handler = get_handler(args)
|
||||
|
||||
# mixup for imagenet
|
||||
mixup_fn = None
|
||||
if args.task in ["imagenet", "in1k"]:
|
||||
mixup_active = args.mixup > 0 or args.cutmix > 0. or args.cutmix_minmax is not None
|
||||
if mixup_active:
|
||||
print("Mixup is activated!")
|
||||
mixup_fn = Mixup(
|
||||
mixup_alpha=args.mixup, cutmix_alpha=args.cutmix, cutmix_minmax=args.cutmix_minmax,
|
||||
prob=args.mixup_prob, switch_prob=args.mixup_switch_prob, mode=args.mixup_mode,
|
||||
label_smoothing=args.label_smoothing, num_classes=args.nb_classes)
|
||||
|
||||
if args.eval:
|
||||
data_loader_test = create_downstream_dataset(args, is_eval=True)
|
||||
if args.task in ["nlvr2", "flickr30k", "coco_retrieval", "imagenet"]:
|
||||
ext_test_stats, task_key = evaluate(data_loader_test, model, device, task_handler)
|
||||
print(f"Accuracy of the network on the {len(data_loader_test.dataset)} test images: {ext_test_stats[task_key]:.3f}%")
|
||||
exit(0)
|
||||
elif args.task == "vqav2":
|
||||
result, _ = evaluate(data_loader_test, model, device, task_handler)
|
||||
utils.dump_predictions(args, result, "vqav2_test")
|
||||
exit(0)
|
||||
elif args.task in ["coco_captioning", "nocaps"]:
|
||||
predictions, _ = evaluate(data_loader_test, model, device, task_handler)
|
||||
prediction_file = utils.dump_predictions(args, predictions, "{}_test".format(args.task))
|
||||
if utils.is_main_process() and args.task == "coco_captioning":
|
||||
captioning_result = utils.coco_caption_eval(args.output_dir, prediction_file, "{}_test".format(args.task))
|
||||
result_file = os.path.join(args.output_dir, f"{args.task}_result.json")
|
||||
print(json.dumps(captioning_result))
|
||||
utils.write_result_to_jsonl(captioning_result, result_file)
|
||||
exit(0)
|
||||
|
||||
print(f"Start training for {args.epochs} epochs")
|
||||
start_time = time.time()
|
||||
|
||||
max_accuracy = 0.0
|
||||
for epoch in range(args.start_epoch, args.epochs):
|
||||
if args.distributed:
|
||||
data_loader_train.sampler.set_epoch(epoch)
|
||||
if log_writer is not None:
|
||||
log_writer.set_step(epoch * num_training_steps_per_epoch * args.update_freq)
|
||||
train_stats = train_one_epoch(
|
||||
model, data_loader_train, optimizer, device, task_handler, epoch,
|
||||
epoch * num_training_steps_per_epoch, lr_schedule_values, loss_scaler,
|
||||
args.clip_grad, args.update_freq, model_ema, log_writer, args.task, mixup_fn,
|
||||
)
|
||||
if args.output_dir and args.save_ckpt:
|
||||
if (epoch + 1) % args.save_ckpt_freq == 0 or epoch + 1 == args.epochs:
|
||||
utils.save_model(
|
||||
args=args, model=model, model_without_ddp=model_without_ddp, optimizer=optimizer,
|
||||
loss_scaler=loss_scaler, epoch=epoch, model_ema=model_ema)
|
||||
if data_loader_val is not None:
|
||||
if args.task not in ["coco_captioning", "nocaps"]:
|
||||
test_stats, task_key = evaluate(data_loader_val, model, device, task_handler)
|
||||
else:
|
||||
predictions, _ = evaluate(data_loader_val, model, device, task_handler)
|
||||
prediction_file = utils.dump_predictions(args, predictions, f"{args.task}_val_e{epoch}")
|
||||
result_file = os.path.join(args.output_dir, f"{args.task}_result_val_e{epoch}.json")
|
||||
task_key = "CIDEr"
|
||||
if utils.is_main_process():
|
||||
test_stats = utils.coco_caption_eval(args.output_dir, prediction_file, "{}_val".format(args.task))
|
||||
utils.write_result_to_jsonl(test_stats, result_file)
|
||||
torch.distributed.barrier()
|
||||
if not utils.is_main_process():
|
||||
test_stats = utils.read_result_from_jsonl(result_file)
|
||||
|
||||
print(f"Performance of the network on the {len(data_loader_val.dataset)} val images: {test_stats[task_key]:.1f}%")
|
||||
if max_accuracy < test_stats[task_key]:
|
||||
max_accuracy = test_stats[task_key]
|
||||
if args.output_dir and args.save_ckpt:
|
||||
utils.save_model(
|
||||
args=args, model=model, model_without_ddp=model_without_ddp, optimizer=optimizer,
|
||||
loss_scaler=loss_scaler, epoch="best", model_ema=model_ema)
|
||||
|
||||
print(f'Max performance: {max_accuracy:.2f}%')
|
||||
if log_writer is not None:
|
||||
log_writer.update(acc=test_stats[task_key], head="perf", step=epoch)
|
||||
|
||||
log_stats = {**{f'train_{k}': v for k, v in train_stats.items()},
|
||||
**{f'val_{k}': v for k, v in test_stats.items()},
|
||||
'epoch': epoch,
|
||||
'n_parameters': n_parameters}
|
||||
else:
|
||||
log_stats = {**{f'train_{k}': v for k, v in train_stats.items()},
|
||||
# **{f'test_{k}': v for k, v in test_stats.items()},
|
||||
'epoch': epoch,
|
||||
'n_parameters': n_parameters}
|
||||
|
||||
if args.output_dir and utils.is_main_process():
|
||||
if log_writer is not None:
|
||||
log_writer.flush()
|
||||
with open(os.path.join(args.output_dir, "log.txt"), mode="a", encoding="utf-8") as f:
|
||||
f.write(json.dumps(log_stats) + "\n")
|
||||
|
||||
total_time = time.time() - start_time
|
||||
total_time_str = str(datetime.timedelta(seconds=int(total_time)))
|
||||
print('Training time {}'.format(total_time_str))
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
opts, ds_init = get_args()
|
||||
if opts.output_dir:
|
||||
Path(opts.output_dir).mkdir(parents=True, exist_ok=True)
|
||||
main(opts, ds_init)
|
||||
@@ -0,0 +1,913 @@
|
||||
# --------------------------------------------------------
|
||||
# Image as a Foreign Language: BEiT Pretraining for Vision and Vision-Language Tasks (https://arxiv.org/abs/2208.10442)
|
||||
# Github source: https://github.com/microsoft/unilm/tree/master/beit3
|
||||
# Copyright (c) 2023 Microsoft
|
||||
# Licensed under The MIT License [see LICENSE for details]
|
||||
# --------------------------------------------------------'
|
||||
|
||||
import datetime
|
||||
import io
|
||||
import os
|
||||
import math
|
||||
import time
|
||||
import json
|
||||
import argparse
|
||||
import numpy as np
|
||||
from pathlib import Path
|
||||
from collections import defaultdict, deque
|
||||
from timm.utils import get_state_dict
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch._six import inf
|
||||
from torchmetrics import Metric
|
||||
from tensorboardX import SummaryWriter
|
||||
|
||||
|
||||
def bool_flag(s):
|
||||
"""
|
||||
Parse boolean arguments from the command line.
|
||||
"""
|
||||
FALSY_STRINGS = {"off", "false", "0"}
|
||||
TRUTHY_STRINGS = {"on", "true", "1"}
|
||||
if s.lower() in FALSY_STRINGS:
|
||||
return False
|
||||
elif s.lower() in TRUTHY_STRINGS:
|
||||
return True
|
||||
else:
|
||||
raise argparse.ArgumentTypeError("invalid value for a boolean flag")
|
||||
|
||||
|
||||
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))
|
||||
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)
|
||||
|
||||
|
||||
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 v is None:
|
||||
continue
|
||||
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():
|
||||
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):
|
||||
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'
|
||||
log_msg = [
|
||||
header,
|
||||
'[{0' + space_fmt + '}/{1}]',
|
||||
'eta: {eta}',
|
||||
'{meters}',
|
||||
'time: {time}',
|
||||
'data: {data}'
|
||||
]
|
||||
if torch.cuda.is_available():
|
||||
log_msg.append('max mem: {memory:.0f}')
|
||||
log_msg = self.delimiter.join(log_msg)
|
||||
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(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(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('{} Total time: {} ({:.4f} s / it)'.format(
|
||||
header, total_time_str, total_time / len(iterable)))
|
||||
|
||||
|
||||
class TensorboardLogger(object):
|
||||
def __init__(self, log_dir):
|
||||
self.writer = SummaryWriter(logdir=log_dir)
|
||||
self.step = 0
|
||||
|
||||
def set_step(self, step=None):
|
||||
if step is not None:
|
||||
self.step = step
|
||||
else:
|
||||
self.step += 1
|
||||
|
||||
def update(self, head='scalar', step=None, **kwargs):
|
||||
for k, v in kwargs.items():
|
||||
if v is None:
|
||||
continue
|
||||
if isinstance(v, torch.Tensor):
|
||||
v = v.item()
|
||||
assert isinstance(v, (float, int))
|
||||
self.writer.add_scalar(head + "/" + k, v, self.step if step is None else step)
|
||||
|
||||
def flush(self):
|
||||
self.writer.flush()
|
||||
|
||||
|
||||
def _load_checkpoint_for_ema(model_ema, checkpoint):
|
||||
"""
|
||||
Workaround for ModelEma._load_checkpoint to accept an already-loaded object
|
||||
"""
|
||||
mem_file = io.BytesIO()
|
||||
torch.save(checkpoint, mem_file)
|
||||
mem_file.seek(0)
|
||||
model_ema._load_checkpoint(mem_file)
|
||||
|
||||
|
||||
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 _get_rank_env():
|
||||
if "RANK" in os.environ:
|
||||
return int(os.environ["RANK"])
|
||||
else:
|
||||
return int(os.environ['OMPI_COMM_WORLD_RANK'])
|
||||
|
||||
|
||||
def _get_local_rank_env():
|
||||
if "LOCAL_RANK" in os.environ:
|
||||
return int(os.environ["LOCAL_RANK"])
|
||||
else:
|
||||
return int(os.environ['OMPI_COMM_WORLD_LOCAL_RANK'])
|
||||
|
||||
|
||||
def _get_world_size_env():
|
||||
if "WORLD_SIZE" in os.environ:
|
||||
return int(os.environ["WORLD_SIZE"])
|
||||
else:
|
||||
return int(os.environ['OMPI_COMM_WORLD_SIZE'])
|
||||
|
||||
|
||||
# The implementation code is modified from DeiT (https://github.com/facebookresearch/deit.git)
|
||||
def init_distributed_mode(args):
|
||||
if args.dist_on_itp:
|
||||
args.rank = _get_rank_env()
|
||||
args.world_size = _get_world_size_env() # int(os.environ['OMPI_COMM_WORLD_SIZE'])
|
||||
args.gpu = _get_local_rank_env()
|
||||
args.dist_url = "tcp://%s:%s" % (os.environ['MASTER_ADDR'], os.environ['MASTER_PORT'])
|
||||
os.environ['LOCAL_RANK'] = str(args.gpu)
|
||||
os.environ['RANK'] = str(args.rank)
|
||||
os.environ['WORLD_SIZE'] = str(args.world_size)
|
||||
# ["RANK", "WORLD_SIZE", "MASTER_ADDR", "MASTER_PORT", "LOCAL_RANK"]
|
||||
elif 'RANK' in os.environ and 'WORLD_SIZE' in os.environ:
|
||||
args.rank = int(os.environ["RANK"])
|
||||
args.world_size = int(os.environ['WORLD_SIZE'])
|
||||
args.gpu = int(os.environ['LOCAL_RANK'])
|
||||
elif 'SLURM_PROCID' in os.environ:
|
||||
args.rank = int(os.environ['SLURM_PROCID'])
|
||||
args.gpu = args.rank % torch.cuda.device_count()
|
||||
else:
|
||||
print('Not using distributed mode')
|
||||
args.distributed = False
|
||||
return
|
||||
|
||||
args.distributed = True
|
||||
|
||||
torch.cuda.set_device(args.gpu)
|
||||
args.dist_backend = 'nccl'
|
||||
print('| distributed init (rank {}): {}, gpu {}'.format(
|
||||
args.rank, args.dist_url, args.gpu), flush=True)
|
||||
torch.distributed.init_process_group(
|
||||
backend=args.dist_backend, init_method=args.dist_url,
|
||||
world_size=args.world_size, rank=args.rank,
|
||||
timeout=datetime.timedelta(0, 7200)
|
||||
)
|
||||
torch.distributed.barrier()
|
||||
setup_for_distributed(args.rank == 0)
|
||||
|
||||
|
||||
def load_state_dict(model, state_dict, prefix='', ignore_missing="relative_position_index"):
|
||||
missing_keys = []
|
||||
unexpected_keys = []
|
||||
error_msgs = []
|
||||
# copy state_dict so _load_from_state_dict can modify it
|
||||
metadata = getattr(state_dict, '_metadata', None)
|
||||
state_dict = state_dict.copy()
|
||||
if metadata is not None:
|
||||
state_dict._metadata = metadata
|
||||
|
||||
def load(module, prefix=''):
|
||||
local_metadata = {} if metadata is None else metadata.get(
|
||||
prefix[:-1], {})
|
||||
module._load_from_state_dict(
|
||||
state_dict, prefix, local_metadata, True, missing_keys, unexpected_keys, error_msgs)
|
||||
for name, child in module._modules.items():
|
||||
if child is not None:
|
||||
load(child, prefix + name + '.')
|
||||
|
||||
load(model, prefix=prefix)
|
||||
|
||||
warn_missing_keys = []
|
||||
ignore_missing_keys = []
|
||||
for key in missing_keys:
|
||||
keep_flag = True
|
||||
for ignore_key in ignore_missing.split('|'):
|
||||
if ignore_key in key:
|
||||
keep_flag = False
|
||||
break
|
||||
if keep_flag:
|
||||
warn_missing_keys.append(key)
|
||||
else:
|
||||
ignore_missing_keys.append(key)
|
||||
|
||||
missing_keys = warn_missing_keys
|
||||
|
||||
if len(missing_keys) > 0:
|
||||
print("Weights of {} not initialized from pretrained model: {}".format(
|
||||
model.__class__.__name__, missing_keys))
|
||||
if len(unexpected_keys) > 0:
|
||||
print("Weights from pretrained model not used in {}: {}".format(
|
||||
model.__class__.__name__, unexpected_keys))
|
||||
if len(ignore_missing_keys) > 0:
|
||||
print("Ignored weights of {} not initialized from pretrained model: {}".format(
|
||||
model.__class__.__name__, ignore_missing_keys))
|
||||
if len(error_msgs) > 0:
|
||||
print('\n'.join(error_msgs))
|
||||
|
||||
|
||||
class NativeScalerWithGradNormCount:
|
||||
state_dict_key = "amp_scaler"
|
||||
|
||||
def __init__(self):
|
||||
self._scaler = torch.cuda.amp.GradScaler()
|
||||
|
||||
def __call__(self, loss, optimizer, clip_grad=None, parameters=None, create_graph=False, update_grad=True):
|
||||
self._scaler.scale(loss).backward(create_graph=create_graph)
|
||||
if update_grad:
|
||||
if clip_grad is not None:
|
||||
assert parameters is not None
|
||||
self._scaler.unscale_(optimizer) # unscale the gradients of optimizer's assigned params in-place
|
||||
norm = torch.nn.utils.clip_grad_norm_(parameters, clip_grad)
|
||||
else:
|
||||
self._scaler.unscale_(optimizer)
|
||||
norm = get_grad_norm_(parameters)
|
||||
self._scaler.step(optimizer)
|
||||
self._scaler.update()
|
||||
else:
|
||||
norm = None
|
||||
return norm
|
||||
|
||||
def state_dict(self):
|
||||
return self._scaler.state_dict()
|
||||
|
||||
def load_state_dict(self, state_dict):
|
||||
self._scaler.load_state_dict(state_dict)
|
||||
|
||||
|
||||
def get_grad_norm_(parameters, norm_type: float = 2.0) -> torch.Tensor:
|
||||
if isinstance(parameters, torch.Tensor):
|
||||
parameters = [parameters]
|
||||
parameters = [p for p in parameters if p.grad is not None]
|
||||
norm_type = float(norm_type)
|
||||
if len(parameters) == 0:
|
||||
return torch.tensor(0.)
|
||||
device = parameters[0].grad.device
|
||||
if norm_type == inf:
|
||||
total_norm = max(p.grad.detach().abs().max().to(device) for p in parameters)
|
||||
else:
|
||||
total_norm = torch.norm(torch.stack([torch.norm(p.grad.detach(), norm_type).to(device) for p in parameters]), norm_type)
|
||||
return total_norm
|
||||
|
||||
|
||||
def cosine_scheduler(base_value, final_value, epochs, niter_per_ep, warmup_epochs=0,
|
||||
start_warmup_value=0, warmup_steps=-1, sched_type="cos"):
|
||||
warmup_schedule = np.array([])
|
||||
warmup_iters = warmup_epochs * niter_per_ep
|
||||
if warmup_steps > 0:
|
||||
warmup_iters = warmup_steps
|
||||
print("Set warmup steps = %d" % warmup_iters)
|
||||
if warmup_epochs > 0:
|
||||
warmup_schedule = np.linspace(start_warmup_value, base_value, warmup_iters)
|
||||
|
||||
if sched_type == "cos":
|
||||
iters = np.arange(epochs * niter_per_ep - warmup_iters)
|
||||
schedule = np.array([
|
||||
final_value + 0.5 * (base_value - final_value) * (1 + math.cos(math.pi * i / (len(iters)))) for i in iters])
|
||||
elif sched_type == "linear":
|
||||
schedule = np.linspace(base_value, final_value, epochs * niter_per_ep - warmup_iters)
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
|
||||
schedule = np.concatenate((warmup_schedule, schedule))
|
||||
|
||||
assert len(schedule) == epochs * niter_per_ep
|
||||
return schedule
|
||||
|
||||
|
||||
def save_model(args, epoch, model, model_without_ddp, optimizer, loss_scaler, model_ema=None):
|
||||
output_dir = Path(args.output_dir)
|
||||
if loss_scaler is not None:
|
||||
checkpoint_paths = [output_dir / ('checkpoint-%s.pth' % epoch)]
|
||||
for checkpoint_path in checkpoint_paths:
|
||||
to_save = {
|
||||
'model': model_without_ddp.state_dict(),
|
||||
'optimizer': optimizer.state_dict(),
|
||||
'epoch': epoch,
|
||||
'scaler': loss_scaler.state_dict(),
|
||||
'args': args,
|
||||
}
|
||||
|
||||
if model_ema is not None:
|
||||
to_save['model_ema'] = get_state_dict(model_ema)
|
||||
|
||||
save_on_master(to_save, checkpoint_path)
|
||||
else:
|
||||
client_state = {'epoch': epoch, "args": args}
|
||||
if model_ema is not None:
|
||||
client_state['model_ema'] = get_state_dict(model_ema)
|
||||
model.save_checkpoint(save_dir=args.output_dir, tag="checkpoint-%s" % epoch, client_state=client_state)
|
||||
|
||||
|
||||
def auto_load_model(args, model, model_without_ddp, optimizer, loss_scaler, model_ema=None):
|
||||
output_dir = Path(args.output_dir)
|
||||
if loss_scaler is not None:
|
||||
# torch.amp
|
||||
if args.auto_resume and len(args.resume) == 0:
|
||||
import glob
|
||||
all_checkpoints = glob.glob(os.path.join(output_dir, 'checkpoint-*.pth'))
|
||||
latest_ckpt = -1
|
||||
for ckpt in all_checkpoints:
|
||||
t = ckpt.split('-')[-1].split('.')[0]
|
||||
if t.isdigit():
|
||||
latest_ckpt = max(int(t), latest_ckpt)
|
||||
if latest_ckpt >= 0:
|
||||
args.resume = os.path.join(output_dir, 'checkpoint-%d.pth' % latest_ckpt)
|
||||
print("Auto resume checkpoint: %s" % args.resume)
|
||||
|
||||
if args.resume:
|
||||
if args.resume.startswith('https'):
|
||||
checkpoint = torch.hub.load_state_dict_from_url(
|
||||
args.resume, map_location='cpu', check_hash=True)
|
||||
else:
|
||||
checkpoint = torch.load(args.resume, map_location='cpu')
|
||||
model_without_ddp.load_state_dict(checkpoint['model'])
|
||||
print("Resume checkpoint %s" % args.resume)
|
||||
if 'optimizer' in checkpoint and 'epoch' in checkpoint:
|
||||
optimizer.load_state_dict(checkpoint['optimizer'])
|
||||
args.start_epoch = checkpoint['epoch'] + 1
|
||||
if hasattr(args, 'model_ema') and args.model_ema:
|
||||
_load_checkpoint_for_ema(model_ema, checkpoint['model_ema'])
|
||||
if 'scaler' in checkpoint:
|
||||
loss_scaler.load_state_dict(checkpoint['scaler'])
|
||||
print("With optim & sched!")
|
||||
else:
|
||||
# deepspeed, only support '--auto_resume'.
|
||||
if args.auto_resume:
|
||||
import glob
|
||||
all_checkpoints = glob.glob(os.path.join(output_dir, 'checkpoint-*'))
|
||||
latest_ckpt = -1
|
||||
for ckpt in all_checkpoints:
|
||||
t = ckpt.split('-')[-1].split('.')[0]
|
||||
if t.isdigit():
|
||||
latest_ckpt = max(int(t), latest_ckpt)
|
||||
if latest_ckpt >= 0:
|
||||
args.resume = os.path.join(output_dir, 'checkpoint-%d' % latest_ckpt)
|
||||
print("Auto resume checkpoint: %d" % latest_ckpt)
|
||||
_, client_states = model.load_checkpoint(args.output_dir, tag='checkpoint-%d' % latest_ckpt)
|
||||
args.start_epoch = client_states['epoch'] + 1
|
||||
if model_ema is not None:
|
||||
if args.model_ema:
|
||||
_load_checkpoint_for_ema(model_ema, client_states['model_ema'])
|
||||
|
||||
|
||||
# The implementation code is modified from DeiT (https://github.com/facebookresearch/deit.git)
|
||||
def load_model_and_may_interpolate(ckpt_path, model, model_key, model_prefix):
|
||||
if ckpt_path.startswith('https'):
|
||||
checkpoint = torch.hub.load_state_dict_from_url(
|
||||
ckpt_path, map_location='cpu', check_hash=True)
|
||||
else:
|
||||
checkpoint = torch.load(ckpt_path, map_location='cpu')
|
||||
|
||||
print("Load ckpt from %s" % ckpt_path)
|
||||
checkpoint_model = None
|
||||
for model_key in model_key.split('|'):
|
||||
if model_key in checkpoint:
|
||||
checkpoint_model = checkpoint[model_key]
|
||||
print("Load state_dict by model_key = %s" % model_key)
|
||||
break
|
||||
|
||||
if checkpoint_model is None:
|
||||
checkpoint_model = checkpoint
|
||||
|
||||
state_dict = model.state_dict()
|
||||
for k in ['head.weight', 'head.bias']:
|
||||
if k in checkpoint_model and checkpoint_model[k].shape != state_dict[k].shape:
|
||||
print(f"Removing key {k} from pretrained checkpoint")
|
||||
del checkpoint_model[k]
|
||||
|
||||
# interpolate position embedding
|
||||
for pos_embed_key in ("vision_pos_embed", "pos_embed", "beit3.encoder.embed_positions.A.weight"):
|
||||
if pos_embed_key in checkpoint_model:
|
||||
pos_embed_checkpoint = checkpoint_model[pos_embed_key]
|
||||
embedding_size = pos_embed_checkpoint.shape[-1]
|
||||
if pos_embed_key == "beit3.encoder.embed_positions.A.weight":
|
||||
# being consistent with Fairseq, which starts from 2 for position embedding
|
||||
torchscale_model = True
|
||||
num_patches = model.beit3.vision_embed.num_patches
|
||||
num_extra_tokens = model.beit3.vision_embed.num_position_embeddings() + 2 - num_patches
|
||||
else:
|
||||
torchscale_model = False
|
||||
num_patches = model.patch_embed.num_patches
|
||||
num_extra_tokens = getattr(model, pos_embed_key).shape[-2] - num_patches
|
||||
# height (== width) for the checkpoint position embedding
|
||||
orig_size = int((pos_embed_checkpoint.shape[-2] - num_extra_tokens) ** 0.5)
|
||||
# height (== width) for the new position embedding
|
||||
new_size = int(num_patches ** 0.5)
|
||||
# class_token and dist_token are kept unchanged
|
||||
if orig_size != new_size:
|
||||
print("Position interpolate from %dx%d to %dx%d" % (orig_size, orig_size, new_size, new_size))
|
||||
if torchscale_model:
|
||||
extra_tokens = pos_embed_checkpoint[:num_extra_tokens].unsqueeze(0)
|
||||
# only the position tokens are interpolated
|
||||
pos_tokens = pos_embed_checkpoint[num_extra_tokens:]
|
||||
else:
|
||||
extra_tokens = pos_embed_checkpoint[:, :num_extra_tokens]
|
||||
# only the position tokens are interpolated
|
||||
pos_tokens = pos_embed_checkpoint[:, num_extra_tokens:]
|
||||
pos_tokens = pos_tokens.reshape(-1, orig_size, orig_size, embedding_size).permute(0, 3, 1, 2)
|
||||
pos_tokens = torch.nn.functional.interpolate(
|
||||
pos_tokens, size=(new_size, new_size), mode='bicubic', align_corners=False)
|
||||
pos_tokens = pos_tokens.permute(0, 2, 3, 1).flatten(1, 2)
|
||||
new_pos_embed = torch.cat((extra_tokens, pos_tokens), dim=1)
|
||||
if torchscale_model:
|
||||
new_pos_embed = new_pos_embed.squeeze(0)
|
||||
checkpoint_model[pos_embed_key] = new_pos_embed
|
||||
|
||||
load_state_dict(model, checkpoint_model, prefix=model_prefix)
|
||||
|
||||
|
||||
def create_ds_config(args):
|
||||
args.deepspeed_config = os.path.join(args.output_dir, "deepspeed_config.json")
|
||||
with open(args.deepspeed_config, mode="w") as writer:
|
||||
ds_config = {
|
||||
"train_batch_size": args.batch_size * args.update_freq * get_world_size(),
|
||||
"train_micro_batch_size_per_gpu": args.batch_size,
|
||||
"steps_per_print": 1000,
|
||||
"optimizer": {
|
||||
"type": "Adam",
|
||||
"adam_w_mode": True,
|
||||
"params": {
|
||||
"lr": args.lr,
|
||||
"weight_decay": args.weight_decay,
|
||||
"bias_correction": True,
|
||||
"betas": [
|
||||
args.opt_betas[0],
|
||||
args.opt_betas[1]
|
||||
],
|
||||
"eps": args.opt_eps
|
||||
}
|
||||
},
|
||||
"fp16": {
|
||||
"enabled": True,
|
||||
"loss_scale": 0,
|
||||
"initial_scale_power": getattr(args, "initial_scale_power", 12),
|
||||
"loss_scale_window": 1000,
|
||||
"hysteresis": 2,
|
||||
"min_loss_scale": 1
|
||||
},
|
||||
"amp": {
|
||||
"enabled": False,
|
||||
"opt_level": "O2"
|
||||
}
|
||||
}
|
||||
|
||||
if args.clip_grad is not None:
|
||||
ds_config.update({'gradient_clipping': args.clip_grad})
|
||||
|
||||
if args.zero_stage == 1:
|
||||
ds_config.update({"zero_optimization": {"stage": args.zero_stage, "reduce_bucket_size": 5e8}})
|
||||
elif args.zero_stage > 1:
|
||||
raise NotImplementedError()
|
||||
|
||||
writer.write(json.dumps(ds_config, indent=2))
|
||||
|
||||
|
||||
def merge_batch_tensors_by_dict_key(batch):
|
||||
batch_tensors = {}
|
||||
for tensor_key in batch[0]:
|
||||
if isinstance(batch[0][tensor_key], torch.Tensor):
|
||||
batch_tensors[tensor_key] = torch.stack([d[tensor_key] for d in batch])
|
||||
else:
|
||||
batch_tensors[tensor_key] = torch.tensor([d[tensor_key] for d in batch], dtype=torch.long)
|
||||
return batch_tensors
|
||||
|
||||
|
||||
def get_loss_scale_for_deepspeed(model):
|
||||
optimizer = model.optimizer
|
||||
loss_scale = None
|
||||
if hasattr(optimizer, 'loss_scale'):
|
||||
loss_scale = optimizer.loss_scale
|
||||
elif hasattr(optimizer, 'cur_scale'):
|
||||
loss_scale = optimizer.cur_scale
|
||||
return loss_scale
|
||||
|
||||
|
||||
class GatherLayer(torch.autograd.Function):
|
||||
"""
|
||||
Gather tensors from all workers with support for backward propagation:
|
||||
This implementation does not cut the gradients as torch.distributed.all_gather does.
|
||||
"""
|
||||
@staticmethod
|
||||
def forward(ctx, x):
|
||||
output = [torch.zeros_like(x) for _ in range(dist.get_world_size())]
|
||||
dist.all_gather(output, x)
|
||||
return tuple(output)
|
||||
@staticmethod
|
||||
def backward(ctx, *grads):
|
||||
all_gradients = torch.stack(grads)
|
||||
dist.all_reduce(all_gradients)
|
||||
return all_gradients[dist.get_rank()]
|
||||
|
||||
|
||||
def gather_features(
|
||||
image_features,
|
||||
text_features,
|
||||
):
|
||||
gathered_image_features = GatherLayer.apply(image_features)
|
||||
gathered_text_features = GatherLayer.apply(text_features)
|
||||
all_image_features = torch.cat(gathered_image_features)
|
||||
all_text_features = torch.cat(gathered_text_features)
|
||||
|
||||
return all_image_features, all_text_features
|
||||
|
||||
|
||||
# The implementation code is modified from open_clip (https://github.com/mlfoundations/open_clip.git)
|
||||
class ClipLoss(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
cache_labels=False,
|
||||
rank=0,
|
||||
world_size=1,
|
||||
):
|
||||
super().__init__()
|
||||
self.cache_labels = cache_labels
|
||||
self.rank = rank
|
||||
self.world_size = world_size
|
||||
|
||||
# cache state
|
||||
self.prev_num_logits = 0
|
||||
self.labels = {}
|
||||
|
||||
def forward(self, image_features, text_features, logit_scale):
|
||||
device = image_features.device
|
||||
if self.world_size > 1:
|
||||
all_image_features, all_text_features = gather_features(
|
||||
image_features, text_features
|
||||
)
|
||||
|
||||
logits_per_image = logit_scale * image_features @ all_text_features.T
|
||||
logits_per_text = logit_scale * text_features @ all_image_features.T
|
||||
else:
|
||||
logits_per_image = logit_scale * image_features @ text_features.T
|
||||
logits_per_text = logit_scale * text_features @ image_features.T
|
||||
|
||||
# calculated ground-truth and cache if enabled
|
||||
num_logits = logits_per_image.shape[0]
|
||||
if self.prev_num_logits != num_logits or device not in self.labels:
|
||||
labels = torch.arange(num_logits, device=device, dtype=torch.long)
|
||||
if self.world_size > 1:
|
||||
labels = labels + num_logits * self.rank
|
||||
if self.cache_labels:
|
||||
self.labels[device] = labels
|
||||
self.prev_num_logits = num_logits
|
||||
else:
|
||||
labels = self.labels[device]
|
||||
|
||||
total_loss = (
|
||||
F.cross_entropy(logits_per_image, labels) +
|
||||
F.cross_entropy(logits_per_text, labels)
|
||||
) / 2
|
||||
return total_loss, logits_per_image, logits_per_text
|
||||
|
||||
|
||||
def write_result_to_jsonl(test_stats, result_file):
|
||||
with open(result_file, mode="w", encoding="utf-8") as writer:
|
||||
writer.write(json.dumps(test_stats, indent=None))
|
||||
|
||||
|
||||
def read_result_from_jsonl(result_file):
|
||||
with open(result_file, mode="r", encoding="utf-8") as reader:
|
||||
return json.load(reader)
|
||||
|
||||
|
||||
# The implementation code is from ViLT (https://github.com/dandelin/ViLT.git)
|
||||
class VQAScore(Metric):
|
||||
def __init__(self, dist_sync_on_step=False):
|
||||
super().__init__(dist_sync_on_step=dist_sync_on_step)
|
||||
self.add_state("score", default=torch.tensor(0.0), dist_reduce_fx="sum")
|
||||
self.add_state("total", default=torch.tensor(0.0), dist_reduce_fx="sum")
|
||||
|
||||
def update(self, logits, target):
|
||||
logits, target = (
|
||||
logits.detach().float().to(self.score.device),
|
||||
target.detach().float().to(self.score.device),
|
||||
)
|
||||
logits = torch.max(logits, 1)[1]
|
||||
one_hots = torch.zeros(*target.size()).to(target)
|
||||
one_hots.scatter_(1, logits.view(-1, 1), 1)
|
||||
scores = one_hots * target
|
||||
|
||||
self.score += scores.sum()
|
||||
self.total += len(logits)
|
||||
|
||||
def compute(self):
|
||||
return self.score / self.total
|
||||
|
||||
|
||||
class BertCaptioningLoss(nn.Module):
|
||||
def __init__(self, label_smoothing, drop_worst_ratio, drop_worst_after):
|
||||
super().__init__()
|
||||
self.label_smoothing = label_smoothing
|
||||
self.drop_worst_ratio = drop_worst_ratio
|
||||
self.drop_worst_after = drop_worst_after
|
||||
self.log_soft = nn.LogSoftmax(dim=1)
|
||||
self.kl = nn.KLDivLoss(reduction='none')
|
||||
self.iter = 0
|
||||
|
||||
def forward(self, logits, target, iter):
|
||||
eps = self.label_smoothing
|
||||
n_class = logits.size(1)
|
||||
one_hot = torch.zeros_like(logits).scatter(1, target.view(-1, 1), 1)
|
||||
one_hot = one_hot * (1 - eps) + (1 - one_hot) * eps / (n_class - 1)
|
||||
log_prb = self.log_soft(logits)
|
||||
loss = self.kl(log_prb, one_hot).sum(1)
|
||||
|
||||
if self.drop_worst_ratio > 0 and iter > self.drop_worst_after:
|
||||
loss, _ = torch.topk(loss,
|
||||
k=int(loss.shape[0] * (1-self.drop_worst_ratio)),
|
||||
largest=False)
|
||||
loss = loss.mean()
|
||||
|
||||
return loss
|
||||
|
||||
|
||||
class BeamHypotheses(object):
|
||||
def __init__(self, n_hyp, max_length, length_penalty, early_stopping):
|
||||
"""
|
||||
Initialize n-best list of hypotheses.
|
||||
"""
|
||||
self.max_length = max_length - 1 # ignoring bos_token
|
||||
self.length_penalty = length_penalty
|
||||
self.early_stopping = early_stopping
|
||||
self.n_hyp = n_hyp
|
||||
self.hyp = []
|
||||
self.worst_score = 1e9
|
||||
|
||||
def __len__(self):
|
||||
"""
|
||||
Number of hypotheses in the list.
|
||||
"""
|
||||
return len(self.hyp)
|
||||
|
||||
def add(self, hyp, sum_logprobs):
|
||||
"""
|
||||
Add a new hypothesis to the list.
|
||||
"""
|
||||
score = sum_logprobs / len(hyp) ** self.length_penalty
|
||||
if len(self) < self.n_hyp or score > self.worst_score:
|
||||
self.hyp.append((score, hyp))
|
||||
if len(self) > self.n_hyp:
|
||||
sorted_scores = sorted([(s, idx) for idx, (s, _) in enumerate(self.hyp)])
|
||||
del self.hyp[sorted_scores[0][1]]
|
||||
self.worst_score = sorted_scores[1][0]
|
||||
else:
|
||||
self.worst_score = min(score, self.worst_score)
|
||||
|
||||
def is_done(self, best_sum_logprobs):
|
||||
"""
|
||||
If there are enough hypotheses and that none of the hypotheses being generated
|
||||
can become better than the worst one in the heap, then we are done with this sentence.
|
||||
"""
|
||||
if len(self) < self.n_hyp:
|
||||
return False
|
||||
elif self.early_stopping:
|
||||
return True
|
||||
else:
|
||||
return self.worst_score >= best_sum_logprobs / self.max_length ** self.length_penalty
|
||||
|
||||
|
||||
def dump_predictions(args, result, file_suffix):
|
||||
global_rank = get_rank()
|
||||
jsons = None
|
||||
if global_rank >= 0:
|
||||
output_file = os.path.join(args.task_cache_path, f"submit_{global_rank}_{file_suffix}.json")
|
||||
with open(output_file, "w") as fp:
|
||||
json.dump(result, fp, indent=2)
|
||||
torch.distributed.barrier()
|
||||
|
||||
if global_rank == 0:
|
||||
world_size = get_world_size()
|
||||
jsons = []
|
||||
for i in range(world_size):
|
||||
each_file = os.path.join(args.task_cache_path, f"submit_{i}_{file_suffix}.json")
|
||||
with open(each_file, "r") as fp:
|
||||
jsons += json.load(fp)
|
||||
|
||||
new_jsons = []
|
||||
res_dict = dict()
|
||||
if args.task in ["coco_captioning", "nocaps"]:
|
||||
qid_key = "image_id"
|
||||
else:
|
||||
# for VQAv2
|
||||
qid_key = "question_id"
|
||||
for item in jsons:
|
||||
if item[qid_key] in res_dict:
|
||||
continue
|
||||
new_jsons.append(item)
|
||||
res_dict[item[qid_key]] = item
|
||||
jsons = new_jsons
|
||||
|
||||
torch.distributed.barrier()
|
||||
os.remove(output_file)
|
||||
else:
|
||||
jsons = result
|
||||
|
||||
result_file = os.path.join(args.output_dir, f"submit_{file_suffix}.json")
|
||||
if jsons is not None:
|
||||
with open(result_file, "w") as fp:
|
||||
json.dump(jsons, fp, indent=2)
|
||||
print("Infer %d examples into %s" % (len(jsons), result_file))
|
||||
return result_file
|
||||
|
||||
|
||||
# The evaluation code is from BLIP (https://github.com/salesforce/BLIP)
|
||||
# For nocaps, please submit the prediction file to the evaluate server (https://eval.ai/web/challenges/challenge-page/355/overview) to obtain the final results
|
||||
def coco_caption_eval(gt_dir, results_file, split):
|
||||
from pycocotools.coco import COCO
|
||||
from pycocoevalcap.eval import COCOEvalCap
|
||||
from torchvision.datasets.utils import download_url
|
||||
|
||||
urls = {'coco_captioning_val': 'https://storage.googleapis.com/sfr-vision-language-research/datasets/coco_karpathy_val_gt.json',
|
||||
'coco_captioning_test': 'https://storage.googleapis.com/sfr-vision-language-research/datasets/coco_karpathy_test_gt.json',
|
||||
'nocaps_val': 'https://github.com/addf400/files/releases/download/beit3/nocaps_val_gt.json'}
|
||||
filenames = {'coco_captioning_val':'coco_karpathy_val_gt.json',
|
||||
'coco_captioning_test':'coco_karpathy_test_gt.json',
|
||||
'nocaps_val':'nocaps_val_gt.json'}
|
||||
|
||||
download_url(urls[split], gt_dir)
|
||||
annotation_file = os.path.join(gt_dir, filenames[split])
|
||||
|
||||
# create coco object and coco_result object
|
||||
coco = COCO(annotation_file)
|
||||
coco_result = coco.loadRes(results_file)
|
||||
|
||||
# create coco_eval object by taking coco and coco_result
|
||||
coco_eval = COCOEvalCap(coco, coco_result)
|
||||
|
||||
# evaluate results
|
||||
# SPICE will take a few minutes the first time, but speeds up due to caching
|
||||
coco_eval.evaluate()
|
||||
|
||||
res_dict = dict()
|
||||
for metric, score in coco_eval.eval.items():
|
||||
res_dict[metric] = score
|
||||
|
||||
return res_dict
|
||||
@@ -0,0 +1,30 @@
|
||||
[
|
||||
"wall", "building", "sky", "floor", "tree", "ceiling", "road",
|
||||
"bed", "windowpane", "grass", "cabinet", "sidewalk",
|
||||
"person", "earth", "door", "table", "mountain", "plant",
|
||||
"curtain", "chair", "car", "water", "painting", "sofa",
|
||||
"shelf", "house", "sea", "mirror", "rug", "field", "armchair",
|
||||
"seat", "fence", "desk", "rock", "wardrobe", "lamp",
|
||||
"bathtub", "railing", "cushion", "base", "box", "column",
|
||||
"signboard", "chest of drawers", "counter", "sand", "sink",
|
||||
"skyscraper", "fireplace", "refrigerator", "grandstand",
|
||||
"path", "stairs", "runway", "case", "pool table", "pillow",
|
||||
"screen door", "stairway", "river", "bridge", "bookcase",
|
||||
"blind", "coffee table", "toilet", "flower", "book", "hill",
|
||||
"bench", "countertop", "stove", "palm", "kitchen island",
|
||||
"computer", "swivel chair", "boat", "bar", "arcade machine",
|
||||
"hovel", "bus", "towel", "light", "truck", "tower",
|
||||
"chandelier", "awning", "streetlight", "booth",
|
||||
"television receiver", "airplane", "dirt track", "apparel",
|
||||
"pole", "land", "bannister", "escalator", "ottoman", "bottle",
|
||||
"buffet", "poster", "stage", "van", "ship", "fountain",
|
||||
"conveyer belt", "canopy", "washer", "plaything",
|
||||
"swimming pool", "stool", "barrel", "basket", "waterfall",
|
||||
"tent", "bag", "minibike", "cradle", "oven", "ball", "food",
|
||||
"step", "tank", "trade name", "microwave", "pot", "animal",
|
||||
"bicycle", "lake", "dishwasher", "screen", "blanket",
|
||||
"sculpture", "hood", "sconce", "vase", "traffic light",
|
||||
"tray", "ashcan", "fan", "pier", "crt screen", "plate",
|
||||
"monitor", "bulletin board", "shower", "radiator", "glass",
|
||||
"clock", "flag"
|
||||
]
|
||||
@@ -0,0 +1,117 @@
|
||||
# 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.
|
||||
|
||||
from copy import deepcopy
|
||||
from typing import Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch.nn import functional as F
|
||||
from torchvision.transforms.functional import resize # type: ignore
|
||||
from torchvision.transforms.functional import to_pil_image
|
||||
import random
|
||||
|
||||
|
||||
class RandomScale:
|
||||
"""
|
||||
Resizes images to the longest side 'target_length', as well as provides
|
||||
methods for resizing coordinates and boxes. Provides methods for
|
||||
transforming both numpy array and batched torch tensors.
|
||||
"""
|
||||
|
||||
def __init__(self, max_length: int, min_length: int) -> None:
|
||||
self.max_length = max_length
|
||||
self.min_length = min_length
|
||||
|
||||
def apply_image(self, image: np.ndarray) -> np.ndarray:
|
||||
"""
|
||||
Expects a numpy array with shape HxWxC in uint8 format.
|
||||
"""
|
||||
target_size = self.get_preprocess_shape(
|
||||
image.shape[0], image.shape[1], self.max_length, self.min_length
|
||||
)
|
||||
return np.array(resize(to_pil_image(image), target_size))
|
||||
|
||||
def apply_coords(
|
||||
self, coords: np.ndarray, original_size: Tuple[int, ...]
|
||||
) -> np.ndarray:
|
||||
"""
|
||||
Expects a numpy array of length 2 in the final dimension. Requires the
|
||||
original image size in (H, W) format.
|
||||
"""
|
||||
old_h, old_w = original_size
|
||||
new_h, new_w = self.get_preprocess_shape(
|
||||
original_size[0], original_size[1], self.max_length, self.min_length
|
||||
)
|
||||
coords = deepcopy(coords).astype(float)
|
||||
coords[..., 0] = coords[..., 0] * (new_w / old_w)
|
||||
coords[..., 1] = coords[..., 1] * (new_h / old_h)
|
||||
return coords
|
||||
|
||||
def apply_boxes(
|
||||
self, boxes: np.ndarray, original_size: Tuple[int, ...]
|
||||
) -> np.ndarray:
|
||||
"""
|
||||
Expects a numpy array shape Bx4. Requires the original image size
|
||||
in (H, W) format.
|
||||
"""
|
||||
boxes = self.apply_coords(boxes.reshape(-1, 2, 2), original_size)
|
||||
return boxes.reshape(-1, 4)
|
||||
|
||||
def apply_image_torch(self, image: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Expects batched images with shape BxCxHxW and float format. This
|
||||
transformation may not exactly match apply_image. apply_image is
|
||||
the transformation expected by the model.
|
||||
"""
|
||||
# Expects an image in BCHW format. May not exactly match apply_image.
|
||||
target_size = self.get_preprocess_shape(
|
||||
image.shape[0], image.shape[1], self.max_length, self.min_length
|
||||
)
|
||||
return F.interpolate(
|
||||
image, target_size, mode="bilinear", align_corners=False, antialias=True
|
||||
)
|
||||
|
||||
def apply_coords_torch(
|
||||
self, coords: torch.Tensor, original_size: Tuple[int, ...]
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Expects a torch tensor with length 2 in the last dimension. Requires the
|
||||
original image size in (H, W) format.
|
||||
"""
|
||||
old_h, old_w = original_size
|
||||
new_h, new_w = self.get_preprocess_shape(
|
||||
original_size[0], original_size[1], self.max_length, self.min_length
|
||||
)
|
||||
coords = deepcopy(coords).to(torch.float)
|
||||
coords[..., 0] = coords[..., 0] * (new_w / old_w)
|
||||
coords[..., 1] = coords[..., 1] * (new_h / old_h)
|
||||
return coords
|
||||
|
||||
def apply_boxes_torch(
|
||||
self, boxes: torch.Tensor, original_size: Tuple[int, ...]
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Expects a torch tensor with shape Bx4. Requires the original image
|
||||
size in (H, W) format.
|
||||
"""
|
||||
boxes = self.apply_coords_torch(boxes.reshape(-1, 2, 2), original_size)
|
||||
return boxes.reshape(-1, 4)
|
||||
|
||||
@staticmethod
|
||||
def get_preprocess_shape(
|
||||
oldh: int, oldw: int, max_length: int, min_length: int
|
||||
) -> Tuple[int, int]:
|
||||
"""
|
||||
Compute the output size given input size and target long side length.
|
||||
"""
|
||||
max_scale = max_length * 1.0 / max(oldh, oldw)
|
||||
min_scale = min_length * 1.0 / max(oldh, oldw)
|
||||
scale = min_scale + random.random() * (max_scale-min_scale)
|
||||
newh, neww = oldh * scale, oldw * scale
|
||||
neww = int(neww + 0.5)
|
||||
newh = int(newh + 0.5)
|
||||
return (newh, neww)
|
||||
@@ -0,0 +1,90 @@
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
|
||||
def get_mask_from_json(json_path, img):
|
||||
try:
|
||||
with open(json_path, "r") as r:
|
||||
anno = json.loads(r.read())
|
||||
except:
|
||||
with open(json_path, "r", encoding="cp1252") as r:
|
||||
anno = json.loads(r.read())
|
||||
|
||||
inform = anno["shapes"]
|
||||
comments = anno["text"]
|
||||
is_sentence = anno["is_sentence"]
|
||||
|
||||
height, width = img.shape[:2]
|
||||
|
||||
### sort polies by area
|
||||
area_list = []
|
||||
valid_poly_list = []
|
||||
for i in inform:
|
||||
label_id = i["label"]
|
||||
points = i["points"]
|
||||
if "flag" == label_id.lower(): ## meaningless deprecated annotations
|
||||
continue
|
||||
|
||||
tmp_mask = np.zeros((height, width), dtype=np.uint8)
|
||||
cv2.polylines(tmp_mask, np.array([points], dtype=np.int32), True, 1, 1)
|
||||
cv2.fillPoly(tmp_mask, np.array([points], dtype=np.int32), 1)
|
||||
tmp_area = tmp_mask.sum()
|
||||
|
||||
area_list.append(tmp_area)
|
||||
valid_poly_list.append(i)
|
||||
|
||||
### ground-truth mask
|
||||
sort_index = np.argsort(area_list)[::-1].astype(np.int32)
|
||||
sort_index = list(sort_index)
|
||||
sort_inform = []
|
||||
for s_idx in sort_index:
|
||||
sort_inform.append(valid_poly_list[s_idx])
|
||||
|
||||
mask = np.zeros((height, width), dtype=np.uint8)
|
||||
for i in sort_inform:
|
||||
label_id = i["label"]
|
||||
points = i["points"]
|
||||
|
||||
if "ignore" in label_id.lower():
|
||||
label_value = 255 # ignored during evaluation
|
||||
else:
|
||||
label_value = 1 # target
|
||||
|
||||
cv2.polylines(mask, np.array([points], dtype=np.int32), True, label_value, 1)
|
||||
cv2.fillPoly(mask, np.array([points], dtype=np.int32), label_value)
|
||||
|
||||
return mask, comments, is_sentence
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
data_dir = "./train"
|
||||
vis_dir = "./vis"
|
||||
|
||||
if not os.path.exists(vis_dir):
|
||||
os.makedirs(vis_dir)
|
||||
|
||||
json_path_list = sorted(glob.glob(data_dir + "/*.json"))
|
||||
for json_path in json_path_list:
|
||||
img_path = json_path.replace(".json", ".jpg")
|
||||
img = cv2.imread(img_path)[:, :, ::-1]
|
||||
|
||||
# In generated mask, value 1 denotes valid target region, and value 255 stands for region ignored during evaluaiton.
|
||||
mask, comments, is_sentence = get_mask_from_json(json_path, img)
|
||||
|
||||
## visualization. Green for target, and red for ignore.
|
||||
valid_mask = (mask == 1).astype(np.float32)[:, :, None]
|
||||
ignore_mask = (mask == 255).astype(np.float32)[:, :, None]
|
||||
vis_img = img * (1 - valid_mask) * (1 - ignore_mask) + (
|
||||
(np.array([0, 255, 0]) * 0.6 + img * 0.4) * valid_mask
|
||||
+ (np.array([255, 0, 0]) * 0.6 + img * 0.4) * ignore_mask
|
||||
)
|
||||
vis_img = np.concatenate([img, vis_img], 1)
|
||||
vis_path = os.path.join(
|
||||
vis_dir, json_path.split("/")[-1].replace(".json", ".jpg")
|
||||
)
|
||||
cv2.imwrite(vis_path, vis_img[:, :, ::-1])
|
||||
print("Visualization has been saved to: ", vis_path)
|
||||
@@ -0,0 +1,386 @@
|
||||
import glob
|
||||
import os
|
||||
import random
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from pycocotools import mask
|
||||
|
||||
from model.segment_anything.utils.transforms import ResizeLongestSide
|
||||
|
||||
from .data_processing import get_mask_from_json
|
||||
from .refer import REFER
|
||||
from .refer_seg_dataset import ReferSegDataset
|
||||
from .sem_seg_dataset import SemSegDataset
|
||||
from torchvision import transforms
|
||||
import json
|
||||
from PIL import Image
|
||||
|
||||
def collate_fn(
|
||||
batch, tokenizer=None, local_rank=-1
|
||||
):
|
||||
image_path_list = []
|
||||
images_list = []
|
||||
images_evf_list = []
|
||||
masks_list = []
|
||||
label_list = []
|
||||
resize_list = []
|
||||
sampled_classes_list = []
|
||||
offset_list = [0]
|
||||
cnt = 0
|
||||
inferences = []
|
||||
for (
|
||||
image_path,
|
||||
images,
|
||||
images_evf,
|
||||
masks,
|
||||
label,
|
||||
resize,
|
||||
sampled_classes,
|
||||
inference,
|
||||
) in batch:
|
||||
image_path_list.append(image_path)
|
||||
images_list.append(images)
|
||||
images_evf_list.append(images_evf)
|
||||
label_list.append(label)
|
||||
masks_list.append(masks.float())
|
||||
resize_list.append(resize)
|
||||
sampled_classes_list.extend(sampled_classes)
|
||||
cnt += len(sampled_classes)
|
||||
offset_list.append(cnt)
|
||||
inferences.append(inference)
|
||||
|
||||
input_ids = [
|
||||
tokenizer(prompt, return_tensors="pt").input_ids[0]
|
||||
for prompt in sampled_classes_list
|
||||
]
|
||||
|
||||
input_ids = torch.nn.utils.rnn.pad_sequence(
|
||||
input_ids, batch_first=True, padding_value=tokenizer.pad_token_id
|
||||
)
|
||||
attention_masks = input_ids.ne(tokenizer.pad_token_id)
|
||||
|
||||
if inferences[0] == False:
|
||||
truncate_len = tokenizer.model_max_length
|
||||
|
||||
if input_ids.shape[1] > truncate_len:
|
||||
input_ids = input_ids[:, :truncate_len]
|
||||
targets = targets[:, :truncate_len]
|
||||
attention_masks = attention_masks[:, :truncate_len]
|
||||
|
||||
return {
|
||||
"image_paths": image_path_list,
|
||||
"images": torch.stack(images_list, dim=0),
|
||||
"images_evf": torch.stack(images_evf_list, dim=0),
|
||||
"input_ids": input_ids,
|
||||
"attention_masks": attention_masks,
|
||||
"masks_list": masks_list,
|
||||
"label_list": label_list,
|
||||
"resize_list": resize_list,
|
||||
"offset": torch.LongTensor(offset_list),
|
||||
"sampled_classes_list": sampled_classes_list,
|
||||
"inference": inferences[0],
|
||||
}
|
||||
|
||||
|
||||
class HybridDataset(torch.utils.data.Dataset):
|
||||
pixel_mean = torch.Tensor([123.675, 116.28, 103.53]).view(-1, 1, 1)
|
||||
pixel_std = torch.Tensor([58.395, 57.12, 57.375]).view(-1, 1, 1)
|
||||
img_size = 1024
|
||||
ignore_label = 255
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_image_dir,
|
||||
tokenizer,
|
||||
samples_per_epoch=500 * 8 * 2 * 10,
|
||||
precision: str = "fp32",
|
||||
image_size: int = 224,
|
||||
num_classes_per_sample: int = 3,
|
||||
exclude_val=False,
|
||||
dataset="sem_seg||refer_seg",
|
||||
sample_rate=[9, 3, 3, 1],
|
||||
sem_seg_data="ade20k||cocostuff||pascal_part||mapillary",
|
||||
refer_seg_data="refclef||refcoco||refcoco+||refcocog",
|
||||
explanatory=-1,
|
||||
model_type="ori",
|
||||
transform=ResizeLongestSide(1024),
|
||||
):
|
||||
self.transform=transform
|
||||
self.model_type = model_type
|
||||
self.exclude_val = exclude_val
|
||||
self.dataset = dataset
|
||||
self.samples_per_epoch = samples_per_epoch
|
||||
self.explanatory = explanatory
|
||||
self.num_classes_per_sample = num_classes_per_sample
|
||||
sample_rate = np.array(sample_rate)
|
||||
self.sample_rate = sample_rate / sample_rate.sum()
|
||||
|
||||
self.base_image_dir = base_image_dir
|
||||
self.image_size = image_size
|
||||
self.tokenizer = tokenizer
|
||||
self.precision = precision
|
||||
|
||||
self.datasets = dataset.split("||")
|
||||
|
||||
self.all_datasets = []
|
||||
for dataset in self.datasets:
|
||||
if dataset == "sem_seg":
|
||||
self.all_datasets.append(
|
||||
SemSegDataset(
|
||||
base_image_dir,
|
||||
tokenizer,
|
||||
samples_per_epoch,
|
||||
precision,
|
||||
image_size,
|
||||
num_classes_per_sample,
|
||||
exclude_val,
|
||||
sem_seg_data,
|
||||
self.model_type,
|
||||
self.transform
|
||||
)
|
||||
)
|
||||
elif dataset == "refer_seg":
|
||||
self.all_datasets.append(
|
||||
ReferSegDataset(
|
||||
base_image_dir,
|
||||
tokenizer,
|
||||
samples_per_epoch,
|
||||
precision,
|
||||
image_size,
|
||||
num_classes_per_sample,
|
||||
exclude_val,
|
||||
refer_seg_data,
|
||||
self.model_type,
|
||||
self.transform
|
||||
)
|
||||
)
|
||||
|
||||
def __len__(self):
|
||||
return self.samples_per_epoch
|
||||
|
||||
def __getitem__(self, idx):
|
||||
ind = np.random.choice(list(range(len(self.datasets))), p=self.sample_rate)
|
||||
data = self.all_datasets[ind]
|
||||
inference = False
|
||||
return *data[0], inference
|
||||
|
||||
|
||||
def init_ade20k(base_image_dir):
|
||||
with open("utils/ade20k_classes.json", "r") as f:
|
||||
ade20k_classes = json.load(f)
|
||||
ade20k_classes = np.array(ade20k_classes)
|
||||
image_ids = sorted(
|
||||
os.listdir(os.path.join(base_image_dir, "ade20k/images", "validation"))
|
||||
)
|
||||
ade20k_image_ids = []
|
||||
for x in image_ids:
|
||||
if x.endswith(".jpg"):
|
||||
ade20k_image_ids.append(x[:-4])
|
||||
ade20k_images = []
|
||||
for image_id in ade20k_image_ids: # self.descriptions:
|
||||
ade20k_images.append(
|
||||
os.path.join(
|
||||
base_image_dir,
|
||||
"ade20k",
|
||||
"images",
|
||||
"validation",
|
||||
"{}.jpg".format(image_id),
|
||||
)
|
||||
)
|
||||
ade20k_labels = [
|
||||
x.replace(".jpg", ".png").replace("images", "annotations")
|
||||
for x in ade20k_images
|
||||
]
|
||||
print("ade20k: ", len(ade20k_images))
|
||||
return ade20k_classes, ade20k_images, ade20k_labels
|
||||
|
||||
|
||||
class ValDataset(torch.utils.data.Dataset):
|
||||
pixel_mean = torch.Tensor([123.675, 116.28, 103.53]).view(-1, 1, 1)
|
||||
pixel_std = torch.Tensor([58.395, 57.12, 57.375]).view(-1, 1, 1)
|
||||
img_size = 1024
|
||||
ignore_label = 255
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_image_dir,
|
||||
tokenizer,
|
||||
val_dataset,
|
||||
image_size=224,
|
||||
model_type="ori"
|
||||
):
|
||||
self.model_type = model_type
|
||||
self.base_image_dir = base_image_dir
|
||||
splits = val_dataset.split("|")
|
||||
if len(splits) == 3:
|
||||
ds, splitBy, split = splits
|
||||
base_image_dir = os.path.join(base_image_dir, "refer_seg")
|
||||
refer_api = REFER(base_image_dir, ds, splitBy)
|
||||
ref_ids_val = refer_api.getRefIds(split=split)
|
||||
images_ids_val = refer_api.getImgIds(ref_ids=ref_ids_val)
|
||||
refs_val = refer_api.loadRefs(ref_ids=ref_ids_val)
|
||||
refer_seg_ds = {}
|
||||
refer_seg_ds["images"] = []
|
||||
loaded_images = refer_api.loadImgs(image_ids=images_ids_val)
|
||||
for item in loaded_images:
|
||||
item = item.copy()
|
||||
if ds == "refclef":
|
||||
item["file_name"] = os.path.join(
|
||||
base_image_dir, "images/saiapr_tc-12", item["file_name"]
|
||||
)
|
||||
elif ds in ["refcoco", "refcoco+", "refcocog", "grefcoco"]:
|
||||
item["file_name"] = os.path.join(
|
||||
base_image_dir,
|
||||
"images/mscoco/images/train2014",
|
||||
item["file_name"],
|
||||
)
|
||||
refer_seg_ds["images"].append(item)
|
||||
refer_seg_ds["annotations"] = refer_api.Anns # anns_val
|
||||
|
||||
img2refs = {}
|
||||
for ref in refs_val:
|
||||
image_id = ref["image_id"]
|
||||
img2refs[image_id] = img2refs.get(image_id, []) + [
|
||||
ref,
|
||||
]
|
||||
refer_seg_ds["img2refs"] = img2refs
|
||||
self.refer_seg_ds = refer_seg_ds
|
||||
self.data_type = "refer_seg"
|
||||
elif val_dataset=="ade":
|
||||
ds = "ade"
|
||||
self.classes, self.images, self.labels = init_ade20k(base_image_dir)
|
||||
self.data_type = "sem_seg"
|
||||
|
||||
|
||||
self.ds = ds
|
||||
self.tokenizer = tokenizer
|
||||
self.transform = ResizeLongestSide(1024)
|
||||
self.image_preprocessor = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
transforms.Resize((image_size, image_size), interpolation=3),
|
||||
transforms.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5))
|
||||
])
|
||||
def __len__(self):
|
||||
if self.data_type == "refer_seg":
|
||||
return len(self.refer_seg_ds["images"])
|
||||
else:
|
||||
return len(self.images)
|
||||
|
||||
def preprocess(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Normalize pixel values and pad to a square input."""
|
||||
# Normalize colors
|
||||
x = (x - self.pixel_mean) / self.pixel_std
|
||||
|
||||
if self.model_type=="effi" or self.model_type=="sam2":
|
||||
x = F.interpolate(x.unsqueeze(0), (self.img_size, self.img_size), mode="bilinear").squeeze(0)
|
||||
else:
|
||||
# Pad
|
||||
h, w = x.shape[-2:]
|
||||
padh = self.img_size - h
|
||||
padw = self.img_size - w
|
||||
x = F.pad(x, (0, padw, 0, padh))
|
||||
return x
|
||||
|
||||
def __getitem__(self, idx):
|
||||
if self.data_type == "refer_seg":
|
||||
refer_seg_ds = self.refer_seg_ds
|
||||
images = refer_seg_ds["images"]
|
||||
annotations = refer_seg_ds["annotations"]
|
||||
img2refs = refer_seg_ds["img2refs"]
|
||||
|
||||
image_info = images[idx]
|
||||
image_path = image_info["file_name"]
|
||||
image_id = image_info["id"]
|
||||
|
||||
refs = img2refs[image_id]
|
||||
if len(refs) == 0:
|
||||
raise ValueError("image {} has no refs".format(image_id))
|
||||
|
||||
sents = []
|
||||
ann_ids = []
|
||||
for ref in refs:
|
||||
for sent in ref["sentences"]:
|
||||
sents.append(sent["sent"].strip().lower())
|
||||
ann_ids.append(ref["ann_id"])
|
||||
|
||||
sampled_sents = sents
|
||||
sampled_ann_ids = ann_ids
|
||||
image = cv2.imread(image_path)
|
||||
image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
|
||||
is_sentence = False
|
||||
|
||||
elif self.data_type == "sem_seg":
|
||||
image_path = self.images[idx]
|
||||
label_path = self.labels[idx]
|
||||
label = Image.open(label_path)
|
||||
label = np.array(label)
|
||||
label[label == 0] = 255
|
||||
label -= 1
|
||||
label[label == 254] = 255
|
||||
unique_label = np.unique(label).tolist()
|
||||
if 255 in unique_label:
|
||||
unique_label.remove(255)
|
||||
|
||||
sampled_sents = [self.classes[class_id] for class_id in unique_label]
|
||||
|
||||
img = cv2.imread(image_path)
|
||||
image = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
|
||||
class_ids = unique_label
|
||||
label = torch.from_numpy(label).long()
|
||||
masks = []
|
||||
for class_id in class_ids:
|
||||
masks.append(label == class_id)
|
||||
masks = torch.stack(masks, dim=0)
|
||||
|
||||
# preprocess image for evf
|
||||
image_evf = self.image_preprocessor(image)
|
||||
|
||||
# preprocess image for sam
|
||||
image = self.transform.apply_image(image)
|
||||
resize = image.shape[:2]
|
||||
image = self.preprocess(torch.from_numpy(image).permute(2, 0, 1).contiguous())
|
||||
|
||||
if self.data_type == "refer_seg":
|
||||
masks = []
|
||||
for i, ann_id in enumerate(sampled_ann_ids):
|
||||
ann = annotations[ann_id]
|
||||
if len(ann["segmentation"]) == 0 and sampled_sents[i] != "":
|
||||
m = np.zeros((image_info["height"], image_info["width"], 1))
|
||||
else:
|
||||
if type(ann["segmentation"][0]) == list: # polygon
|
||||
rle = mask.frPyObjects(
|
||||
ann["segmentation"],
|
||||
image_info["height"],
|
||||
image_info["width"],
|
||||
)
|
||||
else:
|
||||
rle = ann["segmentation"]
|
||||
for i in range(len(rle)):
|
||||
if not isinstance(rle[i]["counts"], bytes):
|
||||
rle[i]["counts"] = rle[i]["counts"].encode()
|
||||
m = mask.decode(rle)
|
||||
m = np.sum(
|
||||
m, axis=2
|
||||
) # sometimes there are multiple binary map (corresponding to multiple segs)
|
||||
m = m.astype(np.uint8) # convert to np.uint8
|
||||
masks.append(m)
|
||||
|
||||
if not isinstance(masks, torch.Tensor):
|
||||
masks = np.stack(masks, axis=0)
|
||||
masks = torch.from_numpy(masks)
|
||||
labels = torch.ones(masks.shape[1], masks.shape[2]) * self.ignore_label
|
||||
inference = True
|
||||
|
||||
return (
|
||||
image_path,
|
||||
image,
|
||||
image_evf,
|
||||
masks,
|
||||
labels,
|
||||
resize,
|
||||
sampled_sents,
|
||||
inference,
|
||||
)
|
||||
@@ -0,0 +1,193 @@
|
||||
import contextlib
|
||||
import copy
|
||||
import io
|
||||
import logging
|
||||
import os
|
||||
import random
|
||||
|
||||
import numpy as np
|
||||
import pycocotools.mask as mask_util
|
||||
from detectron2.structures import Boxes, BoxMode, PolygonMasks, RotatedBoxes
|
||||
from detectron2.utils.file_io import PathManager
|
||||
from fvcore.common.timer import Timer
|
||||
from PIL import Image
|
||||
|
||||
"""
|
||||
This file contains functions to parse RefCOCO-format annotations into dicts in "Detectron2 format".
|
||||
"""
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
__all__ = ["load_refcoco_json"]
|
||||
|
||||
|
||||
def load_grefcoco_json(
|
||||
refer_root,
|
||||
dataset_name,
|
||||
splitby,
|
||||
split,
|
||||
image_root,
|
||||
extra_annotation_keys=None,
|
||||
extra_refer_keys=None,
|
||||
):
|
||||
if dataset_name == "refcocop":
|
||||
dataset_name = "refcoco+"
|
||||
if dataset_name == "refcoco" or dataset_name == "refcoco+":
|
||||
splitby == "unc"
|
||||
if dataset_name == "refcocog":
|
||||
assert splitby == "umd" or splitby == "google"
|
||||
|
||||
dataset_id = "_".join([dataset_name, splitby, split])
|
||||
|
||||
from .grefer import G_REFER
|
||||
|
||||
logger.info("Loading dataset {} ({}-{}) ...".format(dataset_name, splitby, split))
|
||||
logger.info("Refcoco root: {}".format(refer_root))
|
||||
timer = Timer()
|
||||
refer_root = PathManager.get_local_path(refer_root)
|
||||
with contextlib.redirect_stdout(io.StringIO()):
|
||||
refer_api = G_REFER(data_root=refer_root, dataset=dataset_name, splitBy=splitby)
|
||||
if timer.seconds() > 1:
|
||||
logger.info(
|
||||
"Loading {} takes {:.2f} seconds.".format(dataset_id, timer.seconds())
|
||||
)
|
||||
|
||||
ref_ids = refer_api.getRefIds(split=split)
|
||||
img_ids = refer_api.getImgIds(ref_ids)
|
||||
refs = refer_api.loadRefs(ref_ids)
|
||||
imgs = [refer_api.loadImgs(ref["image_id"])[0] for ref in refs]
|
||||
anns = [refer_api.loadAnns(ref["ann_id"]) for ref in refs]
|
||||
imgs_refs_anns = list(zip(imgs, refs, anns))
|
||||
|
||||
logger.info(
|
||||
"Loaded {} images, {} referring object sets in G_RefCOCO format from {}".format(
|
||||
len(img_ids), len(ref_ids), dataset_id
|
||||
)
|
||||
)
|
||||
|
||||
dataset_dicts = []
|
||||
|
||||
ann_keys = ["iscrowd", "bbox", "category_id"] + (extra_annotation_keys or [])
|
||||
ref_keys = ["raw", "sent_id"] + (extra_refer_keys or [])
|
||||
|
||||
ann_lib = {}
|
||||
|
||||
NT_count = 0
|
||||
MT_count = 0
|
||||
|
||||
for img_dict, ref_dict, anno_dicts in imgs_refs_anns:
|
||||
record = {}
|
||||
record["source"] = "grefcoco"
|
||||
record["file_name"] = os.path.join(image_root, img_dict["file_name"])
|
||||
record["height"] = img_dict["height"]
|
||||
record["width"] = img_dict["width"]
|
||||
image_id = record["image_id"] = img_dict["id"]
|
||||
|
||||
# Check that information of image, ann and ref match each other
|
||||
# This fails only when the data parsing logic or the annotation file is buggy.
|
||||
assert ref_dict["image_id"] == image_id
|
||||
assert ref_dict["split"] == split
|
||||
if not isinstance(ref_dict["ann_id"], list):
|
||||
ref_dict["ann_id"] = [ref_dict["ann_id"]]
|
||||
|
||||
# No target samples
|
||||
if None in anno_dicts:
|
||||
assert anno_dicts == [None]
|
||||
assert ref_dict["ann_id"] == [-1]
|
||||
record["empty"] = True
|
||||
obj = {key: None for key in ann_keys if key in ann_keys}
|
||||
obj["bbox_mode"] = BoxMode.XYWH_ABS
|
||||
obj["empty"] = True
|
||||
obj = [obj]
|
||||
|
||||
# Multi target samples
|
||||
else:
|
||||
record["empty"] = False
|
||||
obj = []
|
||||
for anno_dict in anno_dicts:
|
||||
ann_id = anno_dict["id"]
|
||||
if anno_dict["iscrowd"]:
|
||||
continue
|
||||
assert anno_dict["image_id"] == image_id
|
||||
assert ann_id in ref_dict["ann_id"]
|
||||
|
||||
if ann_id in ann_lib:
|
||||
ann = ann_lib[ann_id]
|
||||
else:
|
||||
ann = {key: anno_dict[key] for key in ann_keys if key in anno_dict}
|
||||
ann["bbox_mode"] = BoxMode.XYWH_ABS
|
||||
ann["empty"] = False
|
||||
|
||||
segm = anno_dict.get("segmentation", None)
|
||||
assert segm # either list[list[float]] or dict(RLE)
|
||||
if isinstance(segm, dict):
|
||||
if isinstance(segm["counts"], list):
|
||||
# convert to compressed RLE
|
||||
segm = mask_util.frPyObjects(segm, *segm["size"])
|
||||
else:
|
||||
# filter out invalid polygons (< 3 points)
|
||||
segm = [
|
||||
poly
|
||||
for poly in segm
|
||||
if len(poly) % 2 == 0 and len(poly) >= 6
|
||||
]
|
||||
if len(segm) == 0:
|
||||
num_instances_without_valid_segmentation += 1
|
||||
continue # ignore this instance
|
||||
ann["segmentation"] = segm
|
||||
ann_lib[ann_id] = ann
|
||||
|
||||
obj.append(ann)
|
||||
|
||||
record["annotations"] = obj
|
||||
|
||||
# Process referring expressions
|
||||
sents = ref_dict["sentences"]
|
||||
for sent in sents:
|
||||
ref_record = record.copy()
|
||||
ref = {key: sent[key] for key in ref_keys if key in sent}
|
||||
ref["ref_id"] = ref_dict["ref_id"]
|
||||
ref_record["sentence"] = ref
|
||||
dataset_dicts.append(ref_record)
|
||||
# if ref_record['empty']:
|
||||
# NT_count += 1
|
||||
# else:
|
||||
# MT_count += 1
|
||||
|
||||
# logger.info("NT samples: %d, MT samples: %d", NT_count, MT_count)
|
||||
|
||||
# Debug mode
|
||||
# return dataset_dicts[:100]
|
||||
|
||||
return dataset_dicts
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
"""
|
||||
Test the COCO json dataset loader.
|
||||
|
||||
Usage:
|
||||
python -m detectron2.data.datasets.coco \
|
||||
path/to/json path/to/image_root dataset_name
|
||||
|
||||
"dataset_name" can be "coco_2014_minival_100", or other
|
||||
pre-registered ones
|
||||
"""
|
||||
import sys
|
||||
|
||||
REFCOCO_PATH = "/mnt/lustre/hhding/code/ReLA/datasets"
|
||||
COCO_TRAIN_2014_IMAGE_ROOT = "/mnt/lustre/hhding/code/ReLA/datasets/images"
|
||||
REFCOCO_DATASET = "grefcoco"
|
||||
REFCOCO_SPLITBY = "unc"
|
||||
REFCOCO_SPLIT = "train"
|
||||
|
||||
|
||||
dicts = load_grefcoco_json(
|
||||
REFCOCO_PATH,
|
||||
REFCOCO_DATASET,
|
||||
REFCOCO_SPLITBY,
|
||||
REFCOCO_SPLIT,
|
||||
COCO_TRAIN_2014_IMAGE_ROOT,
|
||||
)
|
||||
print(1)
|
||||
@@ -0,0 +1,352 @@
|
||||
"""
|
||||
grefer v0.1
|
||||
This interface provides access to gRefCOCO.
|
||||
|
||||
The following API functions are defined:
|
||||
G_REFER - REFER api class
|
||||
getRefIds - get ref ids that satisfy given filter conditions.
|
||||
getAnnIds - get ann ids that satisfy given filter conditions.
|
||||
getImgIds - get image ids that satisfy given filter conditions.
|
||||
getCatIds - get category ids that satisfy given filter conditions.
|
||||
loadRefs - load refs with the specified ref ids.
|
||||
loadAnns - load anns with the specified ann ids.
|
||||
loadImgs - load images with the specified image ids.
|
||||
loadCats - load category names with the specified category ids.
|
||||
getRefBox - get ref's bounding box [x, y, w, h] given the ref_id
|
||||
showRef - show image, segmentation or box of the referred object with the ref
|
||||
getMaskByRef - get mask and area of the referred object given ref or ref ids
|
||||
getMask - get mask and area of the referred object given ref
|
||||
showMask - show mask of the referred object given ref
|
||||
"""
|
||||
|
||||
import itertools
|
||||
import json
|
||||
import os.path as osp
|
||||
import pickle
|
||||
import time
|
||||
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import skimage.io as io
|
||||
from matplotlib.collections import PatchCollection
|
||||
from matplotlib.patches import Polygon, Rectangle
|
||||
from pycocotools import mask
|
||||
|
||||
|
||||
class G_REFER:
|
||||
def __init__(self, data_root, dataset="grefcoco", splitBy="unc"):
|
||||
# provide data_root folder which contains grefcoco
|
||||
print("loading dataset %s into memory..." % dataset)
|
||||
self.ROOT_DIR = osp.abspath(osp.dirname(__file__))
|
||||
self.DATA_DIR = osp.join(data_root, dataset)
|
||||
if dataset in ["grefcoco"]:
|
||||
self.IMAGE_DIR = osp.join(data_root, "images/train2014")
|
||||
else:
|
||||
raise KeyError("No refer dataset is called [%s]" % dataset)
|
||||
|
||||
tic = time.time()
|
||||
|
||||
# load refs from data/dataset/refs(dataset).json
|
||||
self.data = {}
|
||||
self.data["dataset"] = dataset
|
||||
|
||||
ref_file = osp.join(self.DATA_DIR, f"grefs({splitBy}).p")
|
||||
if osp.exists(ref_file):
|
||||
self.data["refs"] = pickle.load(open(ref_file, "rb"), fix_imports=True)
|
||||
else:
|
||||
ref_file = osp.join(self.DATA_DIR, f"grefs({splitBy}).json")
|
||||
if osp.exists(ref_file):
|
||||
self.data["refs"] = json.load(open(ref_file, "rb"))
|
||||
else:
|
||||
raise FileNotFoundError("JSON file not found")
|
||||
|
||||
# load annotations from data/dataset/instances.json
|
||||
instances_file = osp.join(self.DATA_DIR, "instances.json")
|
||||
instances = json.load(open(instances_file, "r"))
|
||||
self.data["images"] = instances["images"]
|
||||
self.data["annotations"] = instances["annotations"]
|
||||
self.data["categories"] = instances["categories"]
|
||||
|
||||
# create index
|
||||
self.createIndex()
|
||||
print("DONE (t=%.2fs)" % (time.time() - tic))
|
||||
|
||||
@staticmethod
|
||||
def _toList(x):
|
||||
return x if isinstance(x, list) else [x]
|
||||
|
||||
@staticmethod
|
||||
def match_any(a, b):
|
||||
a = a if isinstance(a, list) else [a]
|
||||
b = b if isinstance(b, list) else [b]
|
||||
return set(a) & set(b)
|
||||
|
||||
def createIndex(self):
|
||||
# create sets of mapping
|
||||
# 1) Refs: {ref_id: ref}
|
||||
# 2) Anns: {ann_id: ann}
|
||||
# 3) Imgs: {image_id: image}
|
||||
# 4) Cats: {category_id: category_name}
|
||||
# 5) Sents: {sent_id: sent}
|
||||
# 6) imgToRefs: {image_id: refs}
|
||||
# 7) imgToAnns: {image_id: anns}
|
||||
# 8) refToAnn: {ref_id: ann}
|
||||
# 9) annToRef: {ann_id: ref}
|
||||
# 10) catToRefs: {category_id: refs}
|
||||
# 11) sentToRef: {sent_id: ref}
|
||||
# 12) sentToTokens: {sent_id: tokens}
|
||||
print("creating index...")
|
||||
# fetch info from instances
|
||||
Anns, Imgs, Cats, imgToAnns = {}, {}, {}, {}
|
||||
Anns[-1] = None
|
||||
for ann in self.data["annotations"]:
|
||||
Anns[ann["id"]] = ann
|
||||
imgToAnns[ann["image_id"]] = imgToAnns.get(ann["image_id"], []) + [ann]
|
||||
for img in self.data["images"]:
|
||||
Imgs[img["id"]] = img
|
||||
for cat in self.data["categories"]:
|
||||
Cats[cat["id"]] = cat["name"]
|
||||
|
||||
# fetch info from refs
|
||||
Refs, imgToRefs, refToAnn, annToRef, catToRefs = {}, {}, {}, {}, {}
|
||||
Sents, sentToRef, sentToTokens = {}, {}, {}
|
||||
availableSplits = []
|
||||
for ref in self.data["refs"]:
|
||||
# ids
|
||||
ref_id = ref["ref_id"]
|
||||
ann_id = ref["ann_id"]
|
||||
category_id = ref["category_id"]
|
||||
image_id = ref["image_id"]
|
||||
|
||||
if ref["split"] not in availableSplits:
|
||||
availableSplits.append(ref["split"])
|
||||
|
||||
# add mapping related to ref
|
||||
if ref_id in Refs:
|
||||
print("Duplicate ref id")
|
||||
Refs[ref_id] = ref
|
||||
imgToRefs[image_id] = imgToRefs.get(image_id, []) + [ref]
|
||||
|
||||
category_id = self._toList(category_id)
|
||||
added_cats = []
|
||||
for cat in category_id:
|
||||
if cat not in added_cats:
|
||||
added_cats.append(cat)
|
||||
catToRefs[cat] = catToRefs.get(cat, []) + [ref]
|
||||
|
||||
ann_id = self._toList(ann_id)
|
||||
refToAnn[ref_id] = [Anns[ann] for ann in ann_id]
|
||||
for ann_id_n in ann_id:
|
||||
annToRef[ann_id_n] = annToRef.get(ann_id_n, []) + [ref]
|
||||
|
||||
# add mapping of sent
|
||||
for sent in ref["sentences"]:
|
||||
Sents[sent["sent_id"]] = sent
|
||||
sentToRef[sent["sent_id"]] = ref
|
||||
sentToTokens[sent["sent_id"]] = sent["tokens"]
|
||||
|
||||
# create class members
|
||||
self.Refs = Refs
|
||||
self.Anns = Anns
|
||||
self.Imgs = Imgs
|
||||
self.Cats = Cats
|
||||
self.Sents = Sents
|
||||
self.imgToRefs = imgToRefs
|
||||
self.imgToAnns = imgToAnns
|
||||
self.refToAnn = refToAnn
|
||||
self.annToRef = annToRef
|
||||
self.catToRefs = catToRefs
|
||||
self.sentToRef = sentToRef
|
||||
self.sentToTokens = sentToTokens
|
||||
self.availableSplits = availableSplits
|
||||
print("index created.")
|
||||
|
||||
def getRefIds(self, image_ids=[], cat_ids=[], split=[]):
|
||||
image_ids = self._toList(image_ids)
|
||||
cat_ids = self._toList(cat_ids)
|
||||
split = self._toList(split)
|
||||
|
||||
for s in split:
|
||||
if s not in self.availableSplits:
|
||||
raise ValueError(f"Invalid split name: {s}")
|
||||
|
||||
refs = self.data["refs"]
|
||||
|
||||
if len(image_ids) > 0:
|
||||
lists = [self.imgToRefs[image_id] for image_id in image_ids]
|
||||
refs = list(itertools.chain.from_iterable(lists))
|
||||
if len(cat_ids) > 0:
|
||||
refs = [ref for ref in refs if self.match_any(ref["category_id"], cat_ids)]
|
||||
if len(split) > 0:
|
||||
refs = [ref for ref in refs if ref["split"] in split]
|
||||
|
||||
ref_ids = [ref["ref_id"] for ref in refs]
|
||||
return ref_ids
|
||||
|
||||
def getAnnIds(self, image_ids=[], ref_ids=[]):
|
||||
image_ids = self._toList(image_ids)
|
||||
ref_ids = self._toList(ref_ids)
|
||||
|
||||
if any([len(image_ids), len(ref_ids)]):
|
||||
if len(image_ids) > 0:
|
||||
lists = [
|
||||
self.imgToAnns[image_id]
|
||||
for image_id in image_ids
|
||||
if image_id in self.imgToAnns
|
||||
]
|
||||
anns = list(itertools.chain.from_iterable(lists))
|
||||
else:
|
||||
anns = self.data["annotations"]
|
||||
ann_ids = [ann["id"] for ann in anns]
|
||||
if len(ref_ids) > 0:
|
||||
lists = [self.Refs[ref_id]["ann_id"] for ref_id in ref_ids]
|
||||
anns_by_ref_id = list(itertools.chain.from_iterable(lists))
|
||||
ann_ids = list(set(ann_ids).intersection(set(anns_by_ref_id)))
|
||||
else:
|
||||
ann_ids = [ann["id"] for ann in self.data["annotations"]]
|
||||
|
||||
return ann_ids
|
||||
|
||||
def getImgIds(self, ref_ids=[]):
|
||||
ref_ids = self._toList(ref_ids)
|
||||
|
||||
if len(ref_ids) > 0:
|
||||
image_ids = list(set([self.Refs[ref_id]["image_id"] for ref_id in ref_ids]))
|
||||
else:
|
||||
image_ids = self.Imgs.keys()
|
||||
return image_ids
|
||||
|
||||
def getCatIds(self):
|
||||
return self.Cats.keys()
|
||||
|
||||
def loadRefs(self, ref_ids=[]):
|
||||
return [self.Refs[ref_id] for ref_id in self._toList(ref_ids)]
|
||||
|
||||
def loadAnns(self, ann_ids=[]):
|
||||
if isinstance(ann_ids, str):
|
||||
ann_ids = int(ann_ids)
|
||||
return [self.Anns[ann_id] for ann_id in self._toList(ann_ids)]
|
||||
|
||||
def loadImgs(self, image_ids=[]):
|
||||
return [self.Imgs[image_id] for image_id in self._toList(image_ids)]
|
||||
|
||||
def loadCats(self, cat_ids=[]):
|
||||
return [self.Cats[cat_id] for cat_id in self._toList(cat_ids)]
|
||||
|
||||
def getRefBox(self, ref_id):
|
||||
anns = self.refToAnn[ref_id]
|
||||
return [ann["bbox"] for ann in anns] # [x, y, w, h]
|
||||
|
||||
def showRef(self, ref, seg_box="seg"):
|
||||
ax = plt.gca()
|
||||
# show image
|
||||
image = self.Imgs[ref["image_id"]]
|
||||
I = io.imread(osp.join(self.IMAGE_DIR, image["file_name"]))
|
||||
ax.imshow(I)
|
||||
# show refer expression
|
||||
for sid, sent in enumerate(ref["sentences"]):
|
||||
print("%s. %s" % (sid + 1, sent["sent"]))
|
||||
# show segmentations
|
||||
if seg_box == "seg":
|
||||
ann_id = ref["ann_id"]
|
||||
ann = self.Anns[ann_id]
|
||||
polygons = []
|
||||
color = []
|
||||
c = "none"
|
||||
if type(ann["segmentation"][0]) == list:
|
||||
# polygon used for refcoco*
|
||||
for seg in ann["segmentation"]:
|
||||
poly = np.array(seg).reshape((len(seg) / 2, 2))
|
||||
polygons.append(Polygon(poly, True, alpha=0.4))
|
||||
color.append(c)
|
||||
p = PatchCollection(
|
||||
polygons,
|
||||
facecolors=color,
|
||||
edgecolors=(1, 1, 0, 0),
|
||||
linewidths=3,
|
||||
alpha=1,
|
||||
)
|
||||
ax.add_collection(p) # thick yellow polygon
|
||||
p = PatchCollection(
|
||||
polygons,
|
||||
facecolors=color,
|
||||
edgecolors=(1, 0, 0, 0),
|
||||
linewidths=1,
|
||||
alpha=1,
|
||||
)
|
||||
ax.add_collection(p) # thin red polygon
|
||||
else:
|
||||
# mask used for refclef
|
||||
rle = ann["segmentation"]
|
||||
m = mask.decode(rle)
|
||||
img = np.ones((m.shape[0], m.shape[1], 3))
|
||||
color_mask = np.array([2.0, 166.0, 101.0]) / 255
|
||||
for i in range(3):
|
||||
img[:, :, i] = color_mask[i]
|
||||
ax.imshow(np.dstack((img, m * 0.5)))
|
||||
# show bounding-box
|
||||
elif seg_box == "box":
|
||||
ann_id = ref["ann_id"]
|
||||
ann = self.Anns[ann_id]
|
||||
bbox = self.getRefBox(ref["ref_id"])
|
||||
box_plot = Rectangle(
|
||||
(bbox[0], bbox[1]),
|
||||
bbox[2],
|
||||
bbox[3],
|
||||
fill=False,
|
||||
edgecolor="green",
|
||||
linewidth=3,
|
||||
)
|
||||
ax.add_patch(box_plot)
|
||||
|
||||
def getMask(self, ann):
|
||||
if not ann:
|
||||
return None
|
||||
if ann["iscrowd"]:
|
||||
raise ValueError("Crowd object")
|
||||
image = self.Imgs[ann["image_id"]]
|
||||
if type(ann["segmentation"][0]) == list: # polygon
|
||||
rle = mask.frPyObjects(ann["segmentation"], image["height"], image["width"])
|
||||
else:
|
||||
rle = ann["segmentation"]
|
||||
|
||||
m = mask.decode(rle)
|
||||
m = np.sum(
|
||||
m, axis=2
|
||||
) # sometimes there are multiple binary map (corresponding to multiple segs)
|
||||
m = m.astype(np.uint8) # convert to np.uint8
|
||||
# compute area
|
||||
area = sum(mask.area(rle)) # should be close to ann['area']
|
||||
return {"mask": m, "area": area}
|
||||
|
||||
def getMaskByRef(self, ref=None, ref_id=None, merge=False):
|
||||
if not ref and not ref_id:
|
||||
raise ValueError
|
||||
if ref:
|
||||
ann_ids = ref["ann_id"]
|
||||
ref_id = ref["ref_id"]
|
||||
else:
|
||||
ann_ids = self.getAnnIds(ref_ids=ref_id)
|
||||
|
||||
if ann_ids == [-1]:
|
||||
img = self.Imgs[self.Refs[ref_id]["image_id"]]
|
||||
return {
|
||||
"mask": np.zeros([img["height"], img["width"]], dtype=np.uint8),
|
||||
"empty": True,
|
||||
}
|
||||
|
||||
anns = self.loadAnns(ann_ids)
|
||||
mask_list = [self.getMask(ann) for ann in anns if not ann["iscrowd"]]
|
||||
|
||||
if merge:
|
||||
merged_masks = sum([mask["mask"] for mask in mask_list])
|
||||
merged_masks[np.where(merged_masks > 1)] = 1
|
||||
return {"mask": merged_masks, "empty": False}
|
||||
else:
|
||||
return mask_list
|
||||
|
||||
def showMask(self, ref):
|
||||
M = self.getMask(ref)
|
||||
msk = M["mask"]
|
||||
ax = plt.gca()
|
||||
ax.imshow(msk)
|
||||
@@ -0,0 +1,391 @@
|
||||
__author__ = "licheng"
|
||||
|
||||
"""
|
||||
This interface provides access to four datasets:
|
||||
1) refclef
|
||||
2) refcoco
|
||||
3) refcoco+
|
||||
4) refcocog
|
||||
split by unc and google
|
||||
|
||||
The following API functions are defined:
|
||||
REFER - REFER api class
|
||||
getRefIds - get ref ids that satisfy given filter conditions.
|
||||
getAnnIds - get ann ids that satisfy given filter conditions.
|
||||
getImgIds - get image ids that satisfy given filter conditions.
|
||||
getCatIds - get category ids that satisfy given filter conditions.
|
||||
loadRefs - load refs with the specified ref ids.
|
||||
loadAnns - load anns with the specified ann ids.
|
||||
loadImgs - load images with the specified image ids.
|
||||
loadCats - load category names with the specified category ids.
|
||||
getRefBox - get ref's bounding box [x, y, w, h] given the ref_id
|
||||
showRef - show image, segmentation or box of the referred object with the ref
|
||||
getMask - get mask and area of the referred object given ref
|
||||
showMask - show mask of the referred object given ref
|
||||
"""
|
||||
|
||||
import itertools
|
||||
import json
|
||||
import os.path as osp
|
||||
import pickle
|
||||
import sys
|
||||
import time
|
||||
from pprint import pprint
|
||||
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import skimage.io as io
|
||||
from matplotlib.collections import PatchCollection
|
||||
from matplotlib.patches import Polygon, Rectangle
|
||||
from pycocotools import mask
|
||||
|
||||
|
||||
class REFER:
|
||||
def __init__(self, data_root, dataset="refcoco", splitBy="unc"):
|
||||
# provide data_root folder which contains refclef, refcoco, refcoco+ and refcocog
|
||||
# also provide dataset name and splitBy information
|
||||
# e.g., dataset = 'refcoco', splitBy = 'unc'
|
||||
print("loading dataset %s into memory..." % dataset)
|
||||
self.ROOT_DIR = osp.abspath(osp.dirname(__file__))
|
||||
self.DATA_DIR = osp.join(data_root, dataset)
|
||||
if dataset in ["refcoco", "refcoco+", "refcocog"]:
|
||||
self.IMAGE_DIR = osp.join(data_root, "images/mscoco/images/train2014")
|
||||
elif dataset == "refclef":
|
||||
self.IMAGE_DIR = osp.join(data_root, "images/saiapr_tc-12")
|
||||
else:
|
||||
print("No refer dataset is called [%s]" % dataset)
|
||||
sys.exit()
|
||||
|
||||
self.dataset = dataset
|
||||
|
||||
# load refs from data/dataset/refs(dataset).json
|
||||
tic = time.time()
|
||||
|
||||
ref_file = osp.join(self.DATA_DIR, "refs(" + splitBy + ").p")
|
||||
print("ref_file: ", ref_file)
|
||||
self.data = {}
|
||||
self.data["dataset"] = dataset
|
||||
self.data["refs"] = pickle.load(open(ref_file, "rb"))
|
||||
|
||||
# load annotations from data/dataset/instances.json
|
||||
instances_file = osp.join(self.DATA_DIR, "instances.json")
|
||||
instances = json.load(open(instances_file, "rb"))
|
||||
self.data["images"] = instances["images"]
|
||||
self.data["annotations"] = instances["annotations"]
|
||||
self.data["categories"] = instances["categories"]
|
||||
|
||||
# create index
|
||||
self.createIndex()
|
||||
print("DONE (t=%.2fs)" % (time.time() - tic))
|
||||
|
||||
def createIndex(self):
|
||||
# create sets of mapping
|
||||
# 1) Refs: {ref_id: ref}
|
||||
# 2) Anns: {ann_id: ann}
|
||||
# 3) Imgs: {image_id: image}
|
||||
# 4) Cats: {category_id: category_name}
|
||||
# 5) Sents: {sent_id: sent}
|
||||
# 6) imgToRefs: {image_id: refs}
|
||||
# 7) imgToAnns: {image_id: anns}
|
||||
# 8) refToAnn: {ref_id: ann}
|
||||
# 9) annToRef: {ann_id: ref}
|
||||
# 10) catToRefs: {category_id: refs}
|
||||
# 11) sentToRef: {sent_id: ref}
|
||||
# 12) sentToTokens: {sent_id: tokens}
|
||||
print("creating index...")
|
||||
# fetch info from instances
|
||||
Anns, Imgs, Cats, imgToAnns = {}, {}, {}, {}
|
||||
for ann in self.data["annotations"]:
|
||||
Anns[ann["id"]] = ann
|
||||
imgToAnns[ann["image_id"]] = imgToAnns.get(ann["image_id"], []) + [ann]
|
||||
for img in self.data["images"]:
|
||||
Imgs[img["id"]] = img
|
||||
for cat in self.data["categories"]:
|
||||
Cats[cat["id"]] = cat["name"]
|
||||
|
||||
# fetch info from refs
|
||||
Refs, imgToRefs, refToAnn, annToRef, catToRefs = {}, {}, {}, {}, {}
|
||||
Sents, sentToRef, sentToTokens = {}, {}, {}
|
||||
for ref in self.data["refs"]:
|
||||
# ids
|
||||
ref_id = ref["ref_id"]
|
||||
ann_id = ref["ann_id"]
|
||||
category_id = ref["category_id"]
|
||||
image_id = ref["image_id"]
|
||||
|
||||
# add mapping related to ref
|
||||
Refs[ref_id] = ref
|
||||
imgToRefs[image_id] = imgToRefs.get(image_id, []) + [ref]
|
||||
catToRefs[category_id] = catToRefs.get(category_id, []) + [ref]
|
||||
refToAnn[ref_id] = Anns[ann_id]
|
||||
annToRef[ann_id] = ref
|
||||
|
||||
# add mapping of sent
|
||||
for sent in ref["sentences"]:
|
||||
Sents[sent["sent_id"]] = sent
|
||||
sentToRef[sent["sent_id"]] = ref
|
||||
sentToTokens[sent["sent_id"]] = sent["tokens"]
|
||||
|
||||
# create class members
|
||||
self.Refs = Refs
|
||||
self.Anns = Anns
|
||||
self.Imgs = Imgs
|
||||
self.Cats = Cats
|
||||
self.Sents = Sents
|
||||
self.imgToRefs = imgToRefs
|
||||
self.imgToAnns = imgToAnns
|
||||
self.refToAnn = refToAnn
|
||||
self.annToRef = annToRef
|
||||
self.catToRefs = catToRefs
|
||||
self.sentToRef = sentToRef
|
||||
self.sentToTokens = sentToTokens
|
||||
print("index created.")
|
||||
|
||||
def getRefIds(self, image_ids=[], cat_ids=[], ref_ids=[], split=""):
|
||||
image_ids = image_ids if type(image_ids) == list else [image_ids]
|
||||
cat_ids = cat_ids if type(cat_ids) == list else [cat_ids]
|
||||
ref_ids = ref_ids if type(ref_ids) == list else [ref_ids]
|
||||
|
||||
if len(image_ids) == len(cat_ids) == len(ref_ids) == len(split) == 0:
|
||||
refs = self.data["refs"]
|
||||
else:
|
||||
if not len(image_ids) == 0:
|
||||
refs = [self.imgToRefs[image_id] for image_id in image_ids]
|
||||
else:
|
||||
refs = self.data["refs"]
|
||||
if not len(cat_ids) == 0:
|
||||
refs = [ref for ref in refs if ref["category_id"] in cat_ids]
|
||||
if not len(ref_ids) == 0:
|
||||
refs = [ref for ref in refs if ref["ref_id"] in ref_ids]
|
||||
if not len(split) == 0:
|
||||
if split in ["testA", "testB", "testC"]:
|
||||
refs = [
|
||||
ref for ref in refs if split[-1] in ref["split"]
|
||||
] # we also consider testAB, testBC, ...
|
||||
elif split in ["testAB", "testBC", "testAC"]:
|
||||
refs = [
|
||||
ref for ref in refs if ref["split"] == split
|
||||
] # rarely used I guess...
|
||||
elif split == "test":
|
||||
refs = [ref for ref in refs if "test" in ref["split"]]
|
||||
elif split == "train" or split == "val":
|
||||
refs = [ref for ref in refs if ref["split"] == split]
|
||||
else:
|
||||
print("No such split [%s]" % split)
|
||||
sys.exit()
|
||||
ref_ids = [ref["ref_id"] for ref in refs]
|
||||
return ref_ids
|
||||
|
||||
def getAnnIds(self, image_ids=[], cat_ids=[], ref_ids=[]):
|
||||
image_ids = image_ids if type(image_ids) == list else [image_ids]
|
||||
cat_ids = cat_ids if type(cat_ids) == list else [cat_ids]
|
||||
ref_ids = ref_ids if type(ref_ids) == list else [ref_ids]
|
||||
|
||||
if len(image_ids) == len(cat_ids) == len(ref_ids) == 0:
|
||||
ann_ids = [ann["id"] for ann in self.data["annotations"]]
|
||||
else:
|
||||
if not len(image_ids) == 0:
|
||||
lists = [
|
||||
self.imgToAnns[image_id]
|
||||
for image_id in image_ids
|
||||
if image_id in self.imgToAnns
|
||||
] # list of [anns]
|
||||
anns = list(itertools.chain.from_iterable(lists))
|
||||
else:
|
||||
anns = self.data["annotations"]
|
||||
if not len(cat_ids) == 0:
|
||||
anns = [ann for ann in anns if ann["category_id"] in cat_ids]
|
||||
ann_ids = [ann["id"] for ann in anns]
|
||||
if not len(ref_ids) == 0:
|
||||
ids = set(ann_ids).intersection(
|
||||
set([self.Refs[ref_id]["ann_id"] for ref_id in ref_ids])
|
||||
)
|
||||
return ann_ids
|
||||
|
||||
def getImgIds(self, ref_ids=[]):
|
||||
ref_ids = ref_ids if type(ref_ids) == list else [ref_ids]
|
||||
|
||||
if not len(ref_ids) == 0:
|
||||
image_ids = list(set([self.Refs[ref_id]["image_id"] for ref_id in ref_ids]))
|
||||
else:
|
||||
image_ids = self.Imgs.keys()
|
||||
return image_ids
|
||||
|
||||
def getCatIds(self):
|
||||
return self.Cats.keys()
|
||||
|
||||
def loadRefs(self, ref_ids=[]):
|
||||
if type(ref_ids) == list:
|
||||
return [self.Refs[ref_id] for ref_id in ref_ids]
|
||||
elif type(ref_ids) == int:
|
||||
return [self.Refs[ref_ids]]
|
||||
|
||||
def loadAnns(self, ann_ids=[]):
|
||||
if type(ann_ids) == list:
|
||||
return [self.Anns[ann_id] for ann_id in ann_ids]
|
||||
elif type(ann_ids) == int or type(ann_ids) == unicode:
|
||||
return [self.Anns[ann_ids]]
|
||||
|
||||
def loadImgs(self, image_ids=[]):
|
||||
if type(image_ids) == list:
|
||||
return [self.Imgs[image_id] for image_id in image_ids]
|
||||
elif type(image_ids) == int:
|
||||
return [self.Imgs[image_ids]]
|
||||
|
||||
def loadCats(self, cat_ids=[]):
|
||||
if type(cat_ids) == list:
|
||||
return [self.Cats[cat_id] for cat_id in cat_ids]
|
||||
elif type(cat_ids) == int:
|
||||
return [self.Cats[cat_ids]]
|
||||
|
||||
def getRefBox(self, ref_id):
|
||||
ref = self.Refs[ref_id]
|
||||
ann = self.refToAnn[ref_id]
|
||||
return ann["bbox"] # [x, y, w, h]
|
||||
|
||||
def showRef(self, ref, seg_box="seg"):
|
||||
ax = plt.gca()
|
||||
# show image
|
||||
image = self.Imgs[ref["image_id"]]
|
||||
I = io.imread(osp.join(self.IMAGE_DIR, image["file_name"]))
|
||||
ax.imshow(I)
|
||||
# show refer expression
|
||||
for sid, sent in enumerate(ref["sentences"]):
|
||||
print("%s. %s" % (sid + 1, sent["sent"]))
|
||||
# show segmentations
|
||||
if seg_box == "seg":
|
||||
ann_id = ref["ann_id"]
|
||||
ann = self.Anns[ann_id]
|
||||
polygons = []
|
||||
color = []
|
||||
c = "none"
|
||||
if type(ann["segmentation"][0]) == list:
|
||||
# polygon used for refcoco*
|
||||
for seg in ann["segmentation"]:
|
||||
poly = np.array(seg).reshape((len(seg) / 2, 2))
|
||||
polygons.append(Polygon(poly, True, alpha=0.4))
|
||||
color.append(c)
|
||||
p = PatchCollection(
|
||||
polygons,
|
||||
facecolors=color,
|
||||
edgecolors=(1, 1, 0, 0),
|
||||
linewidths=3,
|
||||
alpha=1,
|
||||
)
|
||||
ax.add_collection(p) # thick yellow polygon
|
||||
p = PatchCollection(
|
||||
polygons,
|
||||
facecolors=color,
|
||||
edgecolors=(1, 0, 0, 0),
|
||||
linewidths=1,
|
||||
alpha=1,
|
||||
)
|
||||
ax.add_collection(p) # thin red polygon
|
||||
else:
|
||||
# mask used for refclef
|
||||
rle = ann["segmentation"]
|
||||
m = mask.decode(rle)
|
||||
img = np.ones((m.shape[0], m.shape[1], 3))
|
||||
color_mask = np.array([2.0, 166.0, 101.0]) / 255
|
||||
for i in range(3):
|
||||
img[:, :, i] = color_mask[i]
|
||||
ax.imshow(np.dstack((img, m * 0.5)))
|
||||
# show bounding-box
|
||||
elif seg_box == "box":
|
||||
ann_id = ref["ann_id"]
|
||||
ann = self.Anns[ann_id]
|
||||
bbox = self.getRefBox(ref["ref_id"])
|
||||
box_plot = Rectangle(
|
||||
(bbox[0], bbox[1]),
|
||||
bbox[2],
|
||||
bbox[3],
|
||||
fill=False,
|
||||
edgecolor="green",
|
||||
linewidth=3,
|
||||
)
|
||||
ax.add_patch(box_plot)
|
||||
|
||||
def getMask(self, ref):
|
||||
# return mask, area and mask-center
|
||||
ann = self.refToAnn[ref["ref_id"]]
|
||||
image = self.Imgs[ref["image_id"]]
|
||||
if type(ann["segmentation"][0]) == list: # polygon
|
||||
rle = mask.frPyObjects(ann["segmentation"], image["height"], image["width"])
|
||||
else:
|
||||
rle = ann["segmentation"]
|
||||
m = mask.decode(rle)
|
||||
m = np.sum(
|
||||
m, axis=2
|
||||
) # sometimes there are multiple binary map (corresponding to multiple segs)
|
||||
m = m.astype(np.uint8) # convert to np.uint8
|
||||
# compute area
|
||||
area = sum(mask.area(rle)) # should be close to ann['area']
|
||||
return {"mask": m, "area": area}
|
||||
# # position
|
||||
# position_x = np.mean(np.where(m==1)[1]) # [1] means columns (matlab style) -> x (c style)
|
||||
# position_y = np.mean(np.where(m==1)[0]) # [0] means rows (matlab style) -> y (c style)
|
||||
# # mass position (if there were multiple regions, we use the largest one.)
|
||||
# label_m = label(m, connectivity=m.ndim)
|
||||
# regions = regionprops(label_m)
|
||||
# if len(regions) > 0:
|
||||
# largest_id = np.argmax(np.array([props.filled_area for props in regions]))
|
||||
# largest_props = regions[largest_id]
|
||||
# mass_y, mass_x = largest_props.centroid
|
||||
# else:
|
||||
# mass_x, mass_y = position_x, position_y
|
||||
# # if centroid is not in mask, we find the closest point to it from mask
|
||||
# if m[mass_y, mass_x] != 1:
|
||||
# print('Finding closes mask point ...')
|
||||
# kernel = np.ones((10, 10),np.uint8)
|
||||
# me = cv2.erode(m, kernel, iterations = 1)
|
||||
# points = zip(np.where(me == 1)[0].tolist(), np.where(me == 1)[1].tolist()) # row, col style
|
||||
# points = np.array(points)
|
||||
# dist = np.sum((points - (mass_y, mass_x))**2, axis=1)
|
||||
# id = np.argsort(dist)[0]
|
||||
# mass_y, mass_x = points[id]
|
||||
# # return
|
||||
# return {'mask': m, 'area': area, 'position_x': position_x, 'position_y': position_y, 'mass_x': mass_x, 'mass_y': mass_y}
|
||||
# # show image and mask
|
||||
# I = io.imread(osp.join(self.IMAGE_DIR, image['file_name']))
|
||||
# plt.figure()
|
||||
# plt.imshow(I)
|
||||
# ax = plt.gca()
|
||||
# img = np.ones( (m.shape[0], m.shape[1], 3) )
|
||||
# color_mask = np.array([2.0,166.0,101.0])/255
|
||||
# for i in range(3):
|
||||
# img[:,:,i] = color_mask[i]
|
||||
# ax.imshow(np.dstack( (img, m*0.5) ))
|
||||
# plt.show()
|
||||
|
||||
def showMask(self, ref):
|
||||
M = self.getMask(ref)
|
||||
msk = M["mask"]
|
||||
ax = plt.gca()
|
||||
ax.imshow(msk)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
refer = REFER(dataset="refcocog", splitBy="google")
|
||||
ref_ids = refer.getRefIds()
|
||||
print(len(ref_ids))
|
||||
|
||||
print(len(refer.Imgs))
|
||||
print(len(refer.imgToRefs))
|
||||
|
||||
ref_ids = refer.getRefIds(split="train")
|
||||
print("There are %s training referred objects." % len(ref_ids))
|
||||
|
||||
for ref_id in ref_ids:
|
||||
ref = refer.loadRefs(ref_id)[0]
|
||||
if len(ref["sentences"]) < 2:
|
||||
continue
|
||||
|
||||
pprint(ref)
|
||||
print("The label is %s." % refer.Cats[ref["category_id"]])
|
||||
plt.figure()
|
||||
refer.showRef(ref, seg_box="box")
|
||||
plt.show()
|
||||
|
||||
# plt.figure()
|
||||
# refer.showMask(ref)
|
||||
# plt.show()
|
||||
@@ -0,0 +1,254 @@
|
||||
import os
|
||||
import random
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from pycocotools import mask
|
||||
|
||||
from model.segment_anything.utils.transforms import ResizeLongestSide
|
||||
|
||||
from .grefer import G_REFER
|
||||
from .refer import REFER
|
||||
from torchvision import transforms
|
||||
|
||||
|
||||
class ReferSegDataset(torch.utils.data.Dataset):
|
||||
pixel_mean = torch.Tensor([123.675, 116.28, 103.53]).view(-1, 1, 1)
|
||||
pixel_std = torch.Tensor([58.395, 57.12, 57.375]).view(-1, 1, 1)
|
||||
img_size = 1024
|
||||
ignore_label = 255
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_image_dir,
|
||||
tokenizer,
|
||||
samples_per_epoch=500 * 8 * 2 * 10,
|
||||
precision: str = "fp32",
|
||||
image_size: int = 224,
|
||||
num_classes_per_sample: int = 3,
|
||||
exclude_val=False,
|
||||
refer_seg_data="refclef||refcoco||refcoco+||refcocog",
|
||||
model_type="ori",
|
||||
transform=ResizeLongestSide(1024),
|
||||
):
|
||||
self.model_type = model_type
|
||||
self.exclude_val = exclude_val
|
||||
self.samples_per_epoch = samples_per_epoch
|
||||
self.num_classes_per_sample = num_classes_per_sample
|
||||
|
||||
self.base_image_dir = base_image_dir
|
||||
self.tokenizer = tokenizer
|
||||
self.precision = precision
|
||||
self.transform = transform
|
||||
self.image_preprocessor = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
transforms.Resize((image_size, image_size), interpolation=3),
|
||||
transforms.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5))
|
||||
])
|
||||
|
||||
DATA_DIR = os.path.join(base_image_dir, "refer_seg")
|
||||
self.refer_seg_ds_list = refer_seg_data.split(
|
||||
"||"
|
||||
) # ['refclef', 'refcoco', 'refcoco+', 'refcocog']
|
||||
self.refer_seg_data = {}
|
||||
for ds in self.refer_seg_ds_list:
|
||||
if ds == "refcocog":
|
||||
splitBy = "umd"
|
||||
else:
|
||||
splitBy = "unc"
|
||||
|
||||
if ds == "grefcoco":
|
||||
refer_api = G_REFER(DATA_DIR, ds, splitBy)
|
||||
else:
|
||||
refer_api = REFER(DATA_DIR, ds, splitBy)
|
||||
|
||||
ref_ids_train = refer_api.getRefIds(split="train")
|
||||
images_ids_train = refer_api.getImgIds(ref_ids=ref_ids_train)
|
||||
refs_train = refer_api.loadRefs(ref_ids=ref_ids_train)
|
||||
|
||||
refer_seg_ds = {}
|
||||
refer_seg_ds["images"] = []
|
||||
loaded_images = refer_api.loadImgs(image_ids=images_ids_train)
|
||||
|
||||
for item in loaded_images:
|
||||
item = item.copy()
|
||||
if ds == "refclef":
|
||||
item["file_name"] = os.path.join(
|
||||
DATA_DIR, "images/saiapr_tc-12", item["file_name"]
|
||||
)
|
||||
else:
|
||||
item["file_name"] = os.path.join(
|
||||
DATA_DIR, "images/mscoco/images/train2014", item["file_name"]
|
||||
)
|
||||
refer_seg_ds["images"].append(item)
|
||||
refer_seg_ds["annotations"] = refer_api.Anns # anns_train
|
||||
|
||||
print(
|
||||
"dataset {} (refs {}) (train split) has {} images and {} annotations.".format(
|
||||
ds,
|
||||
splitBy,
|
||||
len(refer_seg_ds["images"]),
|
||||
len(refer_seg_ds["annotations"]),
|
||||
)
|
||||
)
|
||||
|
||||
img2refs = {}
|
||||
for ref in refs_train:
|
||||
image_id = ref["image_id"]
|
||||
img2refs[image_id] = img2refs.get(image_id, []) + [
|
||||
ref,
|
||||
]
|
||||
refer_seg_ds["img2refs"] = img2refs
|
||||
self.refer_seg_data[ds] = refer_seg_ds
|
||||
|
||||
def __len__(self):
|
||||
return self.samples_per_epoch
|
||||
|
||||
def preprocess(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Normalize pixel values and pad to a square input."""
|
||||
if self.model_type=="hq":
|
||||
h, w = x.shape[-2:]
|
||||
padh = self.img_size - h
|
||||
padw = self.img_size - w
|
||||
x = F.pad(x, (0, padw, 0, padh), value=128)
|
||||
|
||||
# Normalize colors
|
||||
x = (x - self.pixel_mean) / self.pixel_std
|
||||
|
||||
if self.model_type=="effi" or self.model_type=="sam2":
|
||||
x = F.interpolate(x.unsqueeze(0), (self.img_size, self.img_size), mode="bilinear").squeeze(0)
|
||||
else:
|
||||
# Pad
|
||||
h, w = x.shape[-2:]
|
||||
padh = self.img_size - h
|
||||
padw = self.img_size - w
|
||||
x = F.pad(x, (0, padw, 0, padh))
|
||||
return x
|
||||
|
||||
def __getitem__(self, idx):
|
||||
ds = random.randint(0, len(self.refer_seg_ds_list) - 1)
|
||||
ds = self.refer_seg_ds_list[ds]
|
||||
refer_seg_ds = self.refer_seg_data[ds]
|
||||
images = refer_seg_ds["images"]
|
||||
annotations = refer_seg_ds["annotations"]
|
||||
img2refs = refer_seg_ds["img2refs"]
|
||||
idx = random.randint(0, len(images) - 1)
|
||||
image_info = images[idx]
|
||||
image_path = image_info["file_name"]
|
||||
image_id = image_info["id"]
|
||||
refs = img2refs[image_id]
|
||||
if len(refs) == 0:
|
||||
return self.__getitem__(0)
|
||||
|
||||
sents = []
|
||||
ann_ids = []
|
||||
for ref in refs:
|
||||
for sent in ref["sentences"]:
|
||||
text = sent["sent"]
|
||||
sents.append(text)
|
||||
ann_ids.append(ref["ann_id"])
|
||||
if len(sents) >= self.num_classes_per_sample:
|
||||
sampled_inds = np.random.choice(
|
||||
list(range(len(sents))), size=self.num_classes_per_sample, replace=False
|
||||
)
|
||||
else:
|
||||
sampled_inds = list(range(len(sents)))
|
||||
sampled_sents = np.vectorize(sents.__getitem__)(sampled_inds).tolist()
|
||||
# sampled_ann_ids = np.vectorize(ann_ids.__getitem__)(sampled_inds).tolist()
|
||||
sampled_ann_ids = [ann_ids[ind] for ind in sampled_inds]
|
||||
sampled_classes = sampled_sents
|
||||
image = cv2.imread(image_path)
|
||||
image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
|
||||
|
||||
# preprocess image for evf
|
||||
image_evf = self.image_preprocessor(image)
|
||||
|
||||
image = self.transform.apply_image(image) # preprocess image for sam
|
||||
resize = image.shape[:2]
|
||||
|
||||
image = self.preprocess(torch.from_numpy(image).permute(2, 0, 1).contiguous())
|
||||
|
||||
flag = False
|
||||
masks = []
|
||||
for ann_id in sampled_ann_ids:
|
||||
if isinstance(ann_id, list):
|
||||
flag = True
|
||||
if -1 in ann_id:
|
||||
assert len(ann_id) == 1
|
||||
m = np.zeros((image_info["height"], image_info["width"])).astype(
|
||||
np.uint8
|
||||
)
|
||||
else:
|
||||
m_final = np.zeros(
|
||||
(image_info["height"], image_info["width"])
|
||||
).astype(np.uint8)
|
||||
for ann_id_i in ann_id:
|
||||
ann = annotations[ann_id_i]
|
||||
|
||||
if len(ann["segmentation"]) == 0:
|
||||
m = np.zeros(
|
||||
(image_info["height"], image_info["width"])
|
||||
).astype(np.uint8)
|
||||
else:
|
||||
if type(ann["segmentation"][0]) == list: # polygon
|
||||
rle = mask.frPyObjects(
|
||||
ann["segmentation"],
|
||||
image_info["height"],
|
||||
image_info["width"],
|
||||
)
|
||||
else:
|
||||
rle = ann["segmentation"]
|
||||
for i in range(len(rle)):
|
||||
if not isinstance(rle[i]["counts"], bytes):
|
||||
rle[i]["counts"] = rle[i]["counts"].encode()
|
||||
m = mask.decode(rle)
|
||||
m = np.sum(
|
||||
m, axis=2
|
||||
) # sometimes there are multiple binary map (corresponding to multiple segs)
|
||||
m = m.astype(np.uint8) # convert to np.uint8
|
||||
m_final = m_final | m
|
||||
m = m_final
|
||||
masks.append(m)
|
||||
continue
|
||||
|
||||
ann = annotations[ann_id]
|
||||
|
||||
if len(ann["segmentation"]) == 0:
|
||||
m = np.zeros((image_info["height"], image_info["width"])).astype(
|
||||
np.uint8
|
||||
)
|
||||
masks.append(m)
|
||||
continue
|
||||
|
||||
if type(ann["segmentation"][0]) == list: # polygon
|
||||
rle = mask.frPyObjects(
|
||||
ann["segmentation"], image_info["height"], image_info["width"]
|
||||
)
|
||||
else:
|
||||
rle = ann["segmentation"]
|
||||
for i in range(len(rle)):
|
||||
if not isinstance(rle[i]["counts"], bytes):
|
||||
rle[i]["counts"] = rle[i]["counts"].encode()
|
||||
m = mask.decode(rle)
|
||||
m = np.sum(
|
||||
m, axis=2
|
||||
) # sometimes there are multiple binary map (corresponding to multiple segs)
|
||||
m = m.astype(np.uint8) # convert to np.uint8
|
||||
masks.append(m)
|
||||
|
||||
masks = np.stack(masks, axis=0)
|
||||
|
||||
masks = torch.from_numpy(masks)
|
||||
label = torch.ones(masks.shape[1], masks.shape[2]) * self.ignore_label
|
||||
|
||||
return (
|
||||
image_path,
|
||||
image,
|
||||
image_evf,
|
||||
masks,
|
||||
label,
|
||||
resize,
|
||||
sampled_classes,
|
||||
)
|
||||
@@ -0,0 +1,290 @@
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from PIL import Image
|
||||
from pycocotools.coco import COCO
|
||||
|
||||
from model.segment_anything.utils.transforms import ResizeLongestSide
|
||||
from torchvision import transforms
|
||||
|
||||
def init_mapillary(base_image_dir):
|
||||
mapillary_data_root = os.path.join(base_image_dir, "mapillary")
|
||||
with open(os.path.join(mapillary_data_root, "config_v2.0.json")) as f:
|
||||
mapillary_classes = json.load(f)["labels"]
|
||||
mapillary_classes = [x["readable"].lower() for x in mapillary_classes]
|
||||
mapillary_classes = np.array(mapillary_classes)
|
||||
mapillary_labels = sorted(
|
||||
glob.glob(
|
||||
os.path.join(mapillary_data_root, "training", "v2.0", "labels", "*.png")
|
||||
)
|
||||
)
|
||||
mapillary_images = [
|
||||
x.replace(".png", ".jpg").replace("v2.0/labels", "images")
|
||||
for x in mapillary_labels
|
||||
]
|
||||
print("mapillary: ", len(mapillary_images))
|
||||
return mapillary_classes, mapillary_images, mapillary_labels
|
||||
|
||||
|
||||
def init_ade20k(base_image_dir):
|
||||
with open("utils/ade20k_classes.json", "r") as f:
|
||||
ade20k_classes = json.load(f)
|
||||
ade20k_classes = np.array(ade20k_classes)
|
||||
image_ids = sorted(
|
||||
os.listdir(os.path.join(base_image_dir, "ade20k/images", "training"))
|
||||
)
|
||||
ade20k_image_ids = []
|
||||
for x in image_ids:
|
||||
if x.endswith(".jpg"):
|
||||
ade20k_image_ids.append(x[:-4])
|
||||
ade20k_images = []
|
||||
for image_id in ade20k_image_ids: # self.descriptions:
|
||||
ade20k_images.append(
|
||||
os.path.join(
|
||||
base_image_dir,
|
||||
"ade20k",
|
||||
"images",
|
||||
"training",
|
||||
"{}.jpg".format(image_id),
|
||||
)
|
||||
)
|
||||
ade20k_labels = [
|
||||
x.replace(".jpg", ".png").replace("images", "annotations")
|
||||
for x in ade20k_images
|
||||
]
|
||||
print("ade20k: ", len(ade20k_images))
|
||||
return ade20k_classes, ade20k_images, ade20k_labels
|
||||
|
||||
def init_paco_lvis(base_image_dir):
|
||||
coco_api_paco_lvis = COCO(
|
||||
os.path.join(
|
||||
base_image_dir, "vlpart", "paco", "annotations", "paco_lvis_v1_train.json"
|
||||
)
|
||||
)
|
||||
all_classes = coco_api_paco_lvis.loadCats(coco_api_paco_lvis.getCatIds())
|
||||
class_map_paco_lvis = {}
|
||||
for cat in all_classes:
|
||||
cat_split = cat["name"].strip().split(":")
|
||||
if len(cat_split) == 1:
|
||||
name = cat_split[0].split("_(")[0]
|
||||
else:
|
||||
assert len(cat_split) == 2
|
||||
obj, part = cat_split
|
||||
obj = obj.split("_(")[0]
|
||||
part = part.split("_(")[0]
|
||||
name = (obj, part)
|
||||
class_map_paco_lvis[cat["id"]] = name
|
||||
img_ids = coco_api_paco_lvis.getImgIds()
|
||||
print("paco_lvis: ", len(img_ids))
|
||||
return class_map_paco_lvis, img_ids, coco_api_paco_lvis
|
||||
|
||||
|
||||
def init_pascal_part(base_image_dir):
|
||||
coco_api_pascal_part = COCO(
|
||||
os.path.join(base_image_dir, "vlpart", "pascal_part", "train.json")
|
||||
)
|
||||
all_classes = coco_api_pascal_part.loadCats(coco_api_pascal_part.getCatIds())
|
||||
class_map_pascal_part = {}
|
||||
for cat in all_classes:
|
||||
cat_main, cat_part = cat["name"].strip().split(":")
|
||||
name = (cat_main, cat_part)
|
||||
class_map_pascal_part[cat["id"]] = name
|
||||
img_ids = coco_api_pascal_part.getImgIds()
|
||||
print("pascal_part: ", len(img_ids))
|
||||
return class_map_pascal_part, img_ids, coco_api_pascal_part
|
||||
|
||||
|
||||
class SemSegDataset(torch.utils.data.Dataset):
|
||||
pixel_mean = torch.Tensor([123.675, 116.28, 103.53]).view(-1, 1, 1)
|
||||
pixel_std = torch.Tensor([58.395, 57.12, 57.375]).view(-1, 1, 1)
|
||||
img_size = 1024
|
||||
ignore_label = 255
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_image_dir,
|
||||
tokenizer,
|
||||
samples_per_epoch=500 * 8 * 2 * 10,
|
||||
precision: str = "fp32",
|
||||
image_size: int = 224,
|
||||
num_classes_per_sample: int = 3,
|
||||
exclude_val=False,
|
||||
sem_seg_data="ade20k||pascal_part||mapillary",
|
||||
model_type="ori",
|
||||
transform=ResizeLongestSide(1024),
|
||||
):
|
||||
self.model_type = model_type
|
||||
self.exclude_val = exclude_val
|
||||
self.samples_per_epoch = samples_per_epoch
|
||||
self.num_classes_per_sample = num_classes_per_sample
|
||||
|
||||
self.base_image_dir = base_image_dir
|
||||
self.tokenizer = tokenizer
|
||||
self.precision = precision
|
||||
self.transform = transform
|
||||
self.image_preprocessor = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
transforms.Resize((image_size, image_size), interpolation=3),
|
||||
transforms.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5))
|
||||
])
|
||||
|
||||
self.data2list = {}
|
||||
self.data2classes = {}
|
||||
|
||||
self.sem_seg_datas = sem_seg_data.split("||")
|
||||
for ds in self.sem_seg_datas:
|
||||
classes, images, labels = eval("init_{}".format(ds))(base_image_dir)
|
||||
self.data2list[ds] = (images, labels)
|
||||
self.data2classes[ds] = classes
|
||||
|
||||
|
||||
def __len__(self):
|
||||
return self.samples_per_epoch
|
||||
|
||||
def preprocess(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Normalize pixel values and pad to a square input."""
|
||||
if self.model_type=="hq":
|
||||
h, w = x.shape[-2:]
|
||||
padh = self.img_size - h
|
||||
padw = self.img_size - w
|
||||
x = F.pad(x, (0, padw, 0, padh), value=128)
|
||||
|
||||
# Normalize colors
|
||||
x = (x - self.pixel_mean) / self.pixel_std
|
||||
|
||||
if self.model_type=="effi" or self.model_type=="sam2":
|
||||
x = F.interpolate(x.unsqueeze(0), (self.img_size, self.img_size), mode="bilinear").squeeze(0)
|
||||
else:
|
||||
# Pad
|
||||
h, w = x.shape[-2:]
|
||||
padh = self.img_size - h
|
||||
padw = self.img_size - w
|
||||
x = F.pad(x, (0, padw, 0, padh))
|
||||
return x
|
||||
|
||||
def __getitem__(self, idx):
|
||||
ds = random.randint(0, len(self.sem_seg_datas) - 1)
|
||||
ds = self.sem_seg_datas[ds]
|
||||
|
||||
if ds in ["pascal_part"]:
|
||||
class_map = self.data2classes[ds]
|
||||
img_ids, coco_api = self.data2list[ds]
|
||||
idx = random.randint(0, len(img_ids) - 1)
|
||||
img_id = img_ids[idx]
|
||||
image_info = coco_api.loadImgs([img_id])[0]
|
||||
file_name = image_info["file_name"]
|
||||
file_name = os.path.join(
|
||||
"VOCdevkit", "VOC2010", "JPEGImages", file_name
|
||||
)
|
||||
image_path = os.path.join(self.base_image_dir, "vlpart", ds, file_name)
|
||||
|
||||
image = cv2.imread(image_path)
|
||||
image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
|
||||
|
||||
# preprocess image for evf
|
||||
image_evf = self.image_preprocessor(image)
|
||||
|
||||
image = self.transform.apply_image(image) # preprocess image for sam
|
||||
resize = image.shape[:2]
|
||||
annIds = coco_api.getAnnIds(imgIds=image_info["id"])
|
||||
anns = coco_api.loadAnns(annIds)
|
||||
if len(anns) == 0:
|
||||
return self.__getitem__(0)
|
||||
if len(anns) >= self.num_classes_per_sample:
|
||||
sampled_anns = np.random.choice(
|
||||
anns, size=self.num_classes_per_sample, replace=False
|
||||
).tolist()
|
||||
else:
|
||||
sampled_anns = anns
|
||||
sampled_classes = []
|
||||
for ann in sampled_anns:
|
||||
sampled_cls = class_map[ann["category_id"]]
|
||||
if isinstance(sampled_cls, tuple):
|
||||
obj, part = sampled_cls
|
||||
if random.random() < 0.5:
|
||||
name = obj + " " + part
|
||||
else:
|
||||
name = "the {} of the {}".format(part, obj)
|
||||
else:
|
||||
name = sampled_cls
|
||||
sampled_classes.append(name)
|
||||
|
||||
elif ds in ["ade20k", "mapillary"]:
|
||||
image, labels = self.data2list[ds]
|
||||
idx = random.randint(0, len(image) - 1)
|
||||
image_path = image[idx]
|
||||
label_path = labels[idx]
|
||||
label = Image.open(label_path)
|
||||
label = np.array(label)
|
||||
if ds == "ade20k":
|
||||
label[label == 0] = 255
|
||||
label -= 1
|
||||
label[label == 254] = 255
|
||||
|
||||
img = cv2.imread(image_path)
|
||||
image = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
|
||||
# preprocess image for evf
|
||||
image_evf = self.image_preprocessor(image)
|
||||
image = self.transform.apply_image(image) # preprocess image for sam
|
||||
resize = image.shape[:2]
|
||||
unique_label = np.unique(label).tolist()
|
||||
if 255 in unique_label:
|
||||
unique_label.remove(255)
|
||||
if len(unique_label) == 0:
|
||||
return self.__getitem__(0)
|
||||
|
||||
classes = [self.data2classes[ds][class_id] for class_id in unique_label]
|
||||
if len(classes) >= self.num_classes_per_sample:
|
||||
sampled_classes = np.random.choice(
|
||||
classes, size=self.num_classes_per_sample, replace=False
|
||||
).tolist()
|
||||
else:
|
||||
sampled_classes = classes
|
||||
|
||||
class_ids = []
|
||||
for sampled_cls in sampled_classes:
|
||||
assert len(sampled_cls.split("||")) == 1
|
||||
|
||||
if ds in ["paco_lvis", "pascal_part"]:
|
||||
continue
|
||||
|
||||
class_id = self.data2classes[ds].tolist().index(sampled_cls)
|
||||
class_ids.append(class_id)
|
||||
|
||||
image = self.preprocess(torch.from_numpy(image).permute(2, 0, 1).contiguous())
|
||||
|
||||
if ds in ["pascal_part"]:
|
||||
masks = []
|
||||
for ann in sampled_anns:
|
||||
try:
|
||||
masks.append(coco_api.annToMask(ann))
|
||||
except Exception as e:
|
||||
print(e)
|
||||
return self.__getitem__(0)
|
||||
|
||||
masks = np.stack(masks, axis=0)
|
||||
masks = torch.from_numpy(masks)
|
||||
label = torch.ones(masks.shape[1], masks.shape[2]) * self.ignore_label
|
||||
|
||||
else:
|
||||
label = torch.from_numpy(label).long()
|
||||
masks = []
|
||||
for class_id in class_ids:
|
||||
masks.append(label == class_id)
|
||||
masks = torch.stack(masks, dim=0)
|
||||
# sampled_classes = ["all "+_ for _ in sampled_classes]
|
||||
return (
|
||||
image_path,
|
||||
image,
|
||||
image_evf,
|
||||
masks,
|
||||
label,
|
||||
resize,
|
||||
sampled_classes,
|
||||
)
|
||||
@@ -0,0 +1,127 @@
|
||||
from enum import Enum
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
IGNORE_INDEX = -100
|
||||
|
||||
class Summary(Enum):
|
||||
NONE = 0
|
||||
AVERAGE = 1
|
||||
SUM = 2
|
||||
COUNT = 3
|
||||
|
||||
|
||||
class AverageMeter(object):
|
||||
"""Computes and stores the average and current value"""
|
||||
|
||||
def __init__(self, name, fmt=":f", summary_type=Summary.AVERAGE):
|
||||
self.name = name
|
||||
self.fmt = fmt
|
||||
self.summary_type = summary_type
|
||||
self.reset()
|
||||
|
||||
def reset(self):
|
||||
self.val = 0
|
||||
self.avg = 0
|
||||
self.sum = 0
|
||||
self.count = 0
|
||||
|
||||
def update(self, val, n=1):
|
||||
self.val = val
|
||||
self.sum += val * n
|
||||
self.count += n
|
||||
self.avg = self.sum / self.count
|
||||
|
||||
def all_reduce(self):
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
if isinstance(self.sum, np.ndarray):
|
||||
total = torch.tensor(
|
||||
self.sum.tolist()
|
||||
+ [
|
||||
self.count,
|
||||
],
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
)
|
||||
else:
|
||||
total = torch.tensor(
|
||||
[self.sum, self.count], dtype=torch.float32, device=device
|
||||
)
|
||||
|
||||
dist.all_reduce(total, dist.ReduceOp.SUM, async_op=False)
|
||||
if total.shape[0] > 2:
|
||||
self.sum, self.count = total[:-1].cpu().numpy(), total[-1].cpu().item()
|
||||
else:
|
||||
self.sum, self.count = total.tolist()
|
||||
self.avg = self.sum / (self.count + 1e-5)
|
||||
|
||||
def __str__(self):
|
||||
fmtstr = "{name} {val" + self.fmt + "} ({avg" + self.fmt + "})"
|
||||
return fmtstr.format(**self.__dict__)
|
||||
|
||||
def summary(self):
|
||||
fmtstr = ""
|
||||
if self.summary_type is Summary.NONE:
|
||||
fmtstr = ""
|
||||
elif self.summary_type is Summary.AVERAGE:
|
||||
fmtstr = "{name} {avg:.3f}"
|
||||
elif self.summary_type is Summary.SUM:
|
||||
fmtstr = "{name} {sum:.3f}"
|
||||
elif self.summary_type is Summary.COUNT:
|
||||
fmtstr = "{name} {count:.3f}"
|
||||
else:
|
||||
raise ValueError("invalid summary type %r" % self.summary_type)
|
||||
|
||||
return fmtstr.format(**self.__dict__)
|
||||
|
||||
|
||||
def intersectionAndUnionGPU(output, target, K, ignore_index=255):
|
||||
# 'K' classes, output and target sizes are N or N * L or N * H * W, each value in range 0 to K - 1.
|
||||
assert output.dim() in [1, 2, 3]
|
||||
assert output.shape == target.shape
|
||||
output = output.view(-1)
|
||||
target = target.view(-1)
|
||||
output[target == ignore_index] = ignore_index
|
||||
intersection = output[output == target]
|
||||
area_intersection = torch.histc(intersection, bins=K, min=0, max=K - 1)
|
||||
area_output = torch.histc(output, bins=K, min=0, max=K - 1)
|
||||
area_target = torch.histc(target, bins=K, min=0, max=K - 1)
|
||||
area_union = area_output + area_target - area_intersection
|
||||
return area_intersection, area_union, area_target
|
||||
|
||||
|
||||
class ProgressMeter(object):
|
||||
def __init__(self, num_batches, meters, prefix=""):
|
||||
self.batch_fmtstr = self._get_batch_fmtstr(num_batches)
|
||||
self.meters = meters
|
||||
self.prefix = prefix
|
||||
|
||||
def display(self, batch):
|
||||
entries = [self.prefix + self.batch_fmtstr.format(batch)]
|
||||
entries += [str(meter) for meter in self.meters]
|
||||
print("\t".join(entries))
|
||||
|
||||
def display_summary(self):
|
||||
entries = [" *"]
|
||||
entries += [meter.summary() for meter in self.meters]
|
||||
print(" ".join(entries))
|
||||
|
||||
def _get_batch_fmtstr(self, num_batches):
|
||||
num_digits = len(str(num_batches // 1))
|
||||
fmt = "{:" + str(num_digits) + "d}"
|
||||
return "[" + fmt + "/" + fmt.format(num_batches) + "]"
|
||||
|
||||
|
||||
def dict_to_cuda(input_dict):
|
||||
for k, v in input_dict.items():
|
||||
if isinstance(input_dict[k], torch.Tensor):
|
||||
input_dict[k] = v.cuda(non_blocking=True)
|
||||
elif (
|
||||
isinstance(input_dict[k], list)
|
||||
and len(input_dict[k]) > 0
|
||||
and isinstance(input_dict[k][0], torch.Tensor)
|
||||
):
|
||||
input_dict[k] = [ele.cuda(non_blocking=True) for ele in v]
|
||||
return input_dict
|
||||
@@ -0,0 +1,112 @@
|
||||
'''
|
||||
推理部分代码来自https://github.com/hustvl/EVF-SAM
|
||||
'''
|
||||
import sys
|
||||
|
||||
from .imagefunc import *
|
||||
sys.path.append(os.path.join(os.path.dirname(__file__), 'evf-sam'))
|
||||
from inference import evf_sam_main
|
||||
class EVF_SAM_Ultra:
|
||||
|
||||
def __init__(self):
|
||||
self.NODE_NAME = 'EVF_SAM Ultra'
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
model_list = ["evf-sam2","evf-sam"]
|
||||
precision_list = ["fp16", "bf16", "fp32"]
|
||||
load_in_bit_list = [16, 8, 4]
|
||||
method_list = ['VITMatte', 'VITMatte(local)', 'PyMatting', 'GuidedFilter', ]
|
||||
device_list = ['cuda', 'cpu']
|
||||
return {"required":
|
||||
{
|
||||
"image": ("IMAGE",),
|
||||
"model": (model_list,),
|
||||
"precision": (precision_list,),
|
||||
"load_in_bit": (load_in_bit_list,),
|
||||
"prompt": ("STRING", {"default": "subject"}),
|
||||
"detail_method": (method_list,),
|
||||
"detail_erode": ("INT", {"default": 6, "min": 1, "max": 255, "step": 1}),
|
||||
"detail_dilate": ("INT", {"default": 4, "min": 1, "max": 255, "step": 1}),
|
||||
"black_point": ("FLOAT", {"default": 0.01, "min": 0.01, "max": 0.98, "step": 0.01, "display": "slider"}),
|
||||
"white_point": ("FLOAT", {"default": 0.99, "min": 0.02, "max": 0.99, "step": 0.01, "display": "slider"}),
|
||||
"process_detail": ("BOOLEAN", {"default": True}),
|
||||
"device": (device_list,),
|
||||
"max_megapixels": ("FLOAT", {"default": 2.0, "min": 1, "max": 999, "step": 0.1}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK",)
|
||||
RETURN_NAMES = ("image", "mask",)
|
||||
FUNCTION = "evf_sam_ultra"
|
||||
CATEGORY = '😺dzNodes/LayerMask'
|
||||
|
||||
def evf_sam_ultra(self, image, model, precision, load_in_bit, prompt,
|
||||
detail_method, detail_erode, detail_dilate, black_point, white_point,
|
||||
process_detail, device, max_megapixels,
|
||||
):
|
||||
|
||||
ret_images = []
|
||||
ret_masks = []
|
||||
|
||||
if detail_method == 'VITMatte(local)':
|
||||
local_files_only = True
|
||||
else:
|
||||
local_files_only = False
|
||||
|
||||
if model == 'evf-sam2':
|
||||
model_type = 'sam2'
|
||||
elif model == 'evf-sam':
|
||||
model_type = 'ori'
|
||||
else:
|
||||
model_type = 'effi'
|
||||
|
||||
model_path = ""
|
||||
model_folder_name = 'EVF-SAM'
|
||||
try:
|
||||
model_path = os.path.join(
|
||||
os.path.normpath(folder_paths.folder_names_and_paths[model_folder_name][0][0]), model)
|
||||
except:
|
||||
pass
|
||||
if not os.path.exists(model_path):
|
||||
model_path = os.path.join(folder_paths.models_dir, model_folder_name, model)
|
||||
|
||||
for i in image:
|
||||
i = torch.unsqueeze(i, 0)
|
||||
orig_image = tensor2pil(i).convert('RGB')
|
||||
sys.path.append(os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
mask_image = evf_sam_main(model_path, model_type, precision, load_in_bit, orig_image, prompt)
|
||||
_mask = pil2tensor(mask_image)
|
||||
|
||||
detail_range = detail_erode + detail_dilate
|
||||
if process_detail:
|
||||
if detail_method == 'GuidedFilter':
|
||||
_mask = guided_filter_alpha(i, _mask, detail_range // 6 + 1)
|
||||
_mask = tensor2pil(histogram_remap(_mask, black_point, white_point))
|
||||
elif detail_method == 'PyMatting':
|
||||
_mask = tensor2pil(mask_edge_detail(i, _mask, detail_range // 8 + 1, black_point, white_point))
|
||||
else:
|
||||
_trimap = generate_VITMatte_trimap(_mask, detail_erode, detail_dilate)
|
||||
_mask = generate_VITMatte(orig_image, _trimap, local_files_only=local_files_only, device=device,
|
||||
max_megapixels=max_megapixels)
|
||||
_mask = tensor2pil(histogram_remap(pil2tensor(_mask), black_point, white_point))
|
||||
else:
|
||||
_mask = mask2image(_mask)
|
||||
|
||||
ret_image = RGB2RGBA(orig_image, _mask.convert('L'))
|
||||
ret_images.append(pil2tensor(ret_image))
|
||||
ret_masks.append(image2mask(_mask))
|
||||
|
||||
log(f"{self.NODE_NAME} Processed {len(ret_images)} image(s).", message_type='finish')
|
||||
return (torch.cat(ret_images, dim=0), torch.cat(ret_masks, dim=0),)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"LayerMask: EVFSAMUltra": EVF_SAM_Ultra
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"LayerMask: EVFSAMUltra": "LayerMask: EVF-SAM Ultra"
|
||||
}
|
||||
|
||||
+2
-2
@@ -1,9 +1,9 @@
|
||||
[project]
|
||||
name = "comfyui_layerstyle"
|
||||
description = "A set of nodes for ComfyUI it generate image like Adobe Photoshop's Layer Style. the Drop Shadow is first completed node, and follow-up work is in progress."
|
||||
version = "1.0.29"
|
||||
version = "1.0.30"
|
||||
license = "MIT"
|
||||
dependencies = ["numpy", "pillow", "torch", "matplotlib", "Scipy", "scikit_image", "opencv-contrib-python", "pymatting", "segment_anything", "timm", "addict", "yapf", "colour-science", "wget", "mediapipe", "loguru", "typer_config", "fastapi", "rich", "google-generativeai", "diffusers", "omegaconf", "tqdm", "transformers", "kornia", "image-reward", "ultralytics", "blend_modes", "blind-watermark", "qrcode", "pyzbar", "transparent-background", "huggingface_hub", "psd-tools"]
|
||||
dependencies = ["numpy", "pillow", "torch", "matplotlib", "Scipy", "scikit_image", "opencv-contrib-python", "pymatting", "segment_anything", "timm", "addict", "yapf", "colour-science", "wget", "mediapipe", "loguru", "typer_config", "fastapi", "rich", "google-generativeai", "diffusers", "omegaconf", "tqdm", "transformers", "kornia", "image-reward", "ultralytics", "blend_modes", "blind-watermark", "qrcode", "pyzbar", "transparent-background", "huggingface_hub", "accelerate", "torchscale", "wandb", "hydra-core", "psd-tools"]
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/chflame163/ComfyUI_LayerStyle"
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
huggingface_hub==0.23.3
|
||||
transformers>=4.38.2
|
||||
huggingface_hub>=0.23.3
|
||||
transformers>=4.44.0
|
||||
protobuf>=4.25.3
|
||||
opencv-contrib-python>=4.9.0.80
|
||||
+6
-2
@@ -21,7 +21,7 @@ google-generativeai
|
||||
diffusers
|
||||
omegaconf
|
||||
tqdm
|
||||
transformers>=4.38.2
|
||||
transformers>=4.44.0
|
||||
kornia
|
||||
image-reward
|
||||
ultralytics
|
||||
@@ -30,5 +30,9 @@ blind-watermark
|
||||
qrcode
|
||||
pyzbar
|
||||
transparent-background
|
||||
huggingface_hub
|
||||
huggingface_hub>=0.23.3
|
||||
accelerate
|
||||
torchscale
|
||||
wandb
|
||||
psd-tools
|
||||
hydra-core
|
||||
Binary file not shown.
Reference in New Issue
Block a user