Add files via upload

This commit is contained in:
Yu Yinda
2024-08-25 23:28:50 +08:00
committed by GitHub
parent eee29196a0
commit 41ca5ee65a
11 changed files with 1540 additions and 0 deletions
+57
View File
@@ -0,0 +1,57 @@
## **SegGPT Usage**
- We release the [SegGPT model](https://huggingface.co/BAAI/SegGPT/blob/main/seggpt_vit_large.pth) and inference code for segmentation everything, as well as some example images and videos.
### Installation
```
git clone https://github.com/baaivision/Painter
cd Painter/SegGPT/SegGPT_inference && wget https://huggingface.co/BAAI/SegGPT/resolve/main/seggpt_vit_large.pth
pip install -r requirements.txt
```
### Usage
Everything in an image with a prompt.
```
python seggpt_inference.py \
--input_image examples/hmbb_2.jpg \
--prompt_image examples/hmbb_1.jpg \
--prompt_target examples/hmbb_1_target.png \
--output_dir ./
```
Everything in an image with multiple prompts.
```
python seggpt_inference.py \
--input_image examples/hmbb_3.jpg \
--prompt_image examples/hmbb_1.jpg examples/hmbb_2.jpg \
--prompt_target examples/hmbb_1_target.png examples/hmbb_2_target.png \
--output_dir ./
```
Everything in a video using a prompt image.
```
python seggpt_inference.py \
--input_video examples/video_1.mp4 \
--prompt_image examples/video_1.jpg \
--prompt_target examples/video_1_target.png \
--output_dir ./
```
Everything in a video using the first frame as the prompt.
```
python seggpt_inference.py \
--input_video examples/video_1.mp4 \
--prompt_target examples/video_1_target.png \
--output_dir ./
```
Processing a long video with prompts from both a target image and the predictions of the previous NUM_FRAMES frames.
```
NUM_FRAMES=4
python seggpt_inference.py \
--input_video examples/video_3.mp4 \
--prompt_target examples/video_3_target.png \
--num_frames $NUM_FRAMES \
--output_dir ./
```
<!-- <div align="center">
<image src="rainbow.gif" width="720px" />
</div> -->
+11
View File
@@ -0,0 +1,11 @@
from .seggpt import SegGPT
NODE_CLASS_MAPPINGS = {
"SegGPT": SegGPT
}
# A dictionary that contains the friendly/humanly readable titles for the nodes
NODE_DISPLAY_NAME_MAPPINGS = {
"SegGPT": "SegGPT Node"
}
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
+535
View File
@@ -0,0 +1,535 @@
from functools import partial
import torch
import torch.nn as nn
import torch.nn.functional as F
##########################
import fvcore.nn.weight_init as weight_init
#from detectron2.layers import CNNBlockBase, get_norm
from fairscale.nn.checkpoint import checkpoint_wrapper
from timm.models.layers import DropPath, trunc_normal_
from timm.models.vision_transformer import Mlp
from .util.vitdet_utils import (
PatchEmbed,
add_decomposed_rel_pos,
get_abs_pos,
window_partition,
window_unpartition,
LayerNorm2D,
)
def get_norm(norm_type, num_features, **kwargs):
if norm_type == "BN":
return nn.BatchNorm2d(num_features, **kwargs)
elif norm_type == "GN":
return nn.GroupNorm(num_groups=32, num_channels=num_features, **kwargs)
elif norm_type == "LN":
return nn.LayerNorm(normalized_shape=[num_features], **kwargs)
elif norm_type == "IN":
return nn.InstanceNorm2d(num_features, **kwargs)
else:
raise ValueError(f"Unknown normalization type: {norm_type}")
class Attention(nn.Module):
"""Multi-head Attention block with relative position embeddings."""
def __init__(
self,
dim,
num_heads=8,
qkv_bias=True,
use_rel_pos=False,
rel_pos_zero_init=True,
input_size=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 (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:
# 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))
if not rel_pos_zero_init:
trunc_normal_(self.rel_pos_h, std=0.02)
trunc_normal_(self.rel_pos_w, std=0.02)
def forward(self, x):
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
class CustomCNNBlock(nn.Module):
def __init__(self, in_channels, out_channels):
super(CustomCNNBlock, self).__init__()
self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1)
self.norm = nn.BatchNorm2d(out_channels)
self.relu = nn.ReLU(inplace=True)
def forward(self, x):
return self.relu(self.norm(self.conv(x)))
class ResBottleneckBlock(CustomCNNBlock):
"""
The standard bottleneck residual block without the last activation layer.
It contains 3 conv layers with kernels 1x1, 3x3, 1x1.
"""
def __init__(
self,
in_channels,
out_channels,
bottleneck_channels,
norm="LN",
act_layer=nn.GELU,
):
"""
Args:
in_channels (int): Number of input channels.
out_channels (int): Number of output channels.
bottleneck_channels (int): number of output channels for the 3x3
"bottleneck" conv layers.
norm (str or callable): normalization for all conv layers.
See :func:`layers.get_norm` for supported format.
act_layer (callable): activation for all conv layers.
"""
super().__init__(in_channels, out_channels, 1)
self.conv1 = nn.Conv2d(in_channels, bottleneck_channels, 1, bias=False)
self.norm1 = get_norm(norm, bottleneck_channels)
self.act1 = act_layer()
self.conv2 = nn.Conv2d(
bottleneck_channels,
bottleneck_channels,
3,
padding=1,
bias=False,
)
self.norm2 = get_norm(norm, bottleneck_channels)
self.act2 = act_layer()
self.conv3 = nn.Conv2d(bottleneck_channels, out_channels, 1, bias=False)
self.norm3 = get_norm(norm, out_channels)
for layer in [self.conv1, self.conv2, self.conv3]:
weight_init.c2_msra_fill(layer)
for layer in [self.norm1, self.norm2]:
layer.weight.data.fill_(1.0)
layer.bias.data.zero_()
# zero init last norm layer.
self.norm3.weight.data.zero_()
self.norm3.bias.data.zero_()
def forward(self, x):
out = x
for layer in self.children():
out = layer(out)
out = x + out
return out
class Block(nn.Module):
"""Transformer blocks with support of window attention and residual propagation blocks"""
def __init__(
self,
dim,
num_heads,
mlp_ratio=4.0,
qkv_bias=True,
drop_path=0.0,
norm_layer=nn.LayerNorm,
act_layer=nn.GELU,
use_rel_pos=False,
rel_pos_zero_init=True,
window_size=0,
use_residual_block=False,
input_size=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.
drop_path (float): Stochastic depth rate.
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 not
use window attention.
use_residual_block (bool): If True, use a residual block after the MLP block.
input_size (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.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()
self.norm2 = norm_layer(dim)
self.mlp = Mlp(in_features=dim, hidden_features=int(dim * mlp_ratio), act_layer=act_layer)
self.window_size = window_size
self.use_residual_block = use_residual_block
if use_residual_block:
# Use a residual block with bottleneck channel as dim // 2
self.residual = ResBottleneckBlock(
in_channels=dim,
out_channels=dim,
bottleneck_channels=dim // 2,
norm="LN",
act_layer=act_layer,
)
def forward(self, x, merge=0):
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))
# feature ensemble
if merge > 0:
prompt, inputs = x.split(x.shape[1] // 2, dim=1)
if merge == 1:
num_prompts = x.shape[0] // 2
inputs = inputs.reshape(2, num_prompts, -1)
inputs = inputs.mean(dim=1, keepdim=True).expand_as(inputs)
inputs = inputs.reshape(*prompt.shape)
else:
inputs = inputs.mean(dim=0, keepdim=True).expand_as(inputs)
x = torch.cat([prompt, inputs], dim=1)
x = shortcut + self.drop_path(x)
x = x + self.drop_path(self.mlp(self.norm2(x)))
if self.use_residual_block:
x = self.residual(x.permute(0, 3, 1, 2)).permute(0, 2, 3, 1)
return x
class SegGPT(nn.Module):
def __init__(
self,
img_size=224,
patch_size=16,
in_chans=3,
embed_dim=1024,
depth=24,
num_heads=16,
mlp_ratio=4.,
qkv_bias=True,
drop_path_rate=0.,
norm_layer=nn.LayerNorm,
act_layer=nn.GELU,
use_abs_pos=True,
use_rel_pos=False,
rel_pos_zero_init=True,
window_size=0,
window_block_indexes=(),
residual_block_indexes=(),
use_act_checkpoint=False,
pretrain_img_size=224,
pretrain_use_cls_token=True,
out_feature="last_feat",
decoder_embed_dim=128,
loss_func="smoothl1",
):
super().__init__()
# --------------------------------------------------------------------------
self.pretrain_use_cls_token = pretrain_use_cls_token
self.patch_size = patch_size
self.patch_embed = PatchEmbed(
kernel_size=(patch_size, patch_size),
stride=(patch_size, patch_size),
in_chans=in_chans,
embed_dim=embed_dim,
)
self.patch_embed.num_patches = (img_size[0] // patch_size) * (img_size[1] // patch_size)
self.mask_token = nn.Parameter(torch.zeros(1, 1, 1, embed_dim))
self.segment_token_x = nn.Parameter(torch.zeros(1, 1, 1, embed_dim))
self.segment_token_y = nn.Parameter(torch.zeros(1, 1, 1, embed_dim))
# token for seg types
self.type_token_cls = nn.Parameter(torch.zeros(1, 1, 1, embed_dim))
self.type_token_ins = nn.Parameter(torch.zeros(1, 1, 1, embed_dim))
if use_abs_pos:
# 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) if pretrain_use_cls_token else num_patches
self.pos_embed = nn.Parameter(torch.zeros(1, num_positions, embed_dim), requires_grad=True)
else:
self.pos_embed = None
# stochastic depth decay rule
dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)]
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,
drop_path=dpr[i],
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 in window_block_indexes else 0,
use_residual_block=i in residual_block_indexes,
input_size=(img_size[0] // patch_size, img_size[1] // patch_size),
)
if use_act_checkpoint:
block = checkpoint_wrapper(block)
self.blocks.append(block)
self._out_feature_channels = {out_feature: embed_dim}
self._out_feature_strides = {out_feature: patch_size}
self._out_features = [out_feature]
if self.pos_embed is not None:
trunc_normal_(self.pos_embed, std=0.02)
self.norm = norm_layer(embed_dim)
# --------------------------------------------------------------------------
# --------------------------------------------------------------------------
self.decoder_embed_dim = decoder_embed_dim
self.decoder_embed = nn.Linear(embed_dim*4, patch_size ** 2 * self.decoder_embed_dim, bias=True) # decoder to patch
self.decoder_pred = nn.Sequential(
nn.Conv2d(self.decoder_embed_dim, self.decoder_embed_dim, kernel_size=3, padding=1, ),
LayerNorm2D(self.decoder_embed_dim),
nn.GELU(),
nn.Conv2d(self.decoder_embed_dim, 3, kernel_size=1, bias=True), # decoder to patch
)
# --------------------------------------------------------------------------
self.loss_func = loss_func
torch.nn.init.normal_(self.mask_token, std=.02)
torch.nn.init.normal_(self.segment_token_x, std=.02)
torch.nn.init.normal_(self.segment_token_y, std=.02)
torch.nn.init.normal_(self.type_token_cls, std=.02)
torch.nn.init.normal_(self.type_token_ins, std=.02)
self.apply(self._init_weights)
def _init_weights(self, m):
if isinstance(m, nn.Linear):
trunc_normal_(m.weight, std=0.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)
@torch.jit.ignore
def no_weight_decay(self):
return {'pos_embed', 'cls_token'}
def patchify(self, imgs):
"""
imgs: (N, 3, H, W)
x: (N, L, patch_size**2 *3)
"""
p = self.patch_size
assert imgs.shape[2] == 2 * imgs.shape[3] and imgs.shape[2] % p == 0
w = imgs.shape[3] // p
h = w * 2
x = imgs.reshape(shape=(imgs.shape[0], 3, h, p, w, p))
x = torch.einsum('nchpwq->nhwpqc', x)
x = x.reshape(shape=(imgs.shape[0], h * w, p**2 * 3))
return x
def unpatchify(self, x):
"""
x: (N, L, patch_size**2 *3)
imgs: (N, 3, H, W)
"""
p = self.patch_size
w = int((x.shape[1]*0.5)**.5)
h = w * 2
assert h * w == x.shape[1]
x = x.reshape(shape=(x.shape[0], h, w, p, p, 3))
x = torch.einsum('nhwpqc->nchpwq', x)
imgs = x.reshape(shape=(x.shape[0], 3, h * p, w * p))
return imgs
def forward_encoder(self, imgs, tgts, bool_masked_pos, seg_type, merge_between_batch=-1):
# embed patches
x = self.patch_embed(imgs)
y = self.patch_embed(tgts)
batch_size, Hp, Wp, _ = x.size()
seq_len = Hp * Wp
mask_token = self.mask_token.expand(batch_size, Hp, Wp, -1)
# replace the masked visual tokens by mask_token
w = bool_masked_pos.unsqueeze(-1).type_as(mask_token).reshape(-1, Hp, Wp, 1)
y = y * (1 - w) + mask_token * w
# add pos embed w/o cls token
x = x + self.segment_token_x
y = y + self.segment_token_y
if self.pos_embed is not None:
x = x + get_abs_pos(
self.pos_embed, self.pretrain_use_cls_token, (x.shape[1], x.shape[2])
)
y = y + get_abs_pos(
self.pos_embed, self.pretrain_use_cls_token, (y.shape[1], y.shape[2])
)
# add type tokens for cls and ins
type_emb = torch.zeros(batch_size, 1, 1, self.type_token_cls.shape[-1]).to(x.device)
type_emb[seg_type==0] = self.type_token_cls
type_emb[seg_type==1] = self.type_token_ins
x = x + type_emb
y = y + type_emb
x = torch.cat((x, y), dim=0)
merge_idx = 2
# apply Transformer blocks
out = []
for idx, blk in enumerate(self.blocks):
merge = 0
if merge_between_batch >= 0 and idx >= merge_between_batch:
merge = 1 if merge_idx >= idx else 2
x = blk(x, merge=merge)
if idx == merge_idx:
x = (x[:x.shape[0]//2] + x[x.shape[0]//2:]) * 0.5
if idx in [5, 11, 17, 23]:
out.append(self.norm(x))
return out
def forward_decoder(self, x):
x = torch.cat(x, dim=-1)
x = self.decoder_embed(x) # BxhxwxC
p = self.patch_size
h, w = x.shape[1], x.shape[2]
x = x.reshape(shape=(x.shape[0], h, w, p, p, self.decoder_embed_dim))
x = torch.einsum('nhwpqc->nchpwq', x)
x = x.reshape(shape=(x.shape[0], -1, h * p, w * p))
x = self.decoder_pred(x) # Bx3xHxW
return x
def forward_loss(self, pred, tgts, mask, valid):
"""
tgts: [N, 3, H, W]
pred: [N, 3, H, W]
mask: [N, L], 0 is keep, 1 is remove,
valid: [N, 3, H, W]
"""
mask = mask[:, :, None].repeat(1, 1, self.patch_size**2 * 3)
mask = self.unpatchify(mask)
mask = mask * valid
target = tgts
if self.loss_func == "l1l2":
loss = ((pred - target).abs() + (pred - target) ** 2.) * 0.5
elif self.loss_func == "l1":
loss = (pred - target).abs()
elif self.loss_func == "l2":
loss = (pred - target) ** 2.
elif self.loss_func == "smoothl1":
loss = F.smooth_l1_loss(pred, target, reduction="none", beta=0.01)
loss = (loss * mask).sum() / mask.sum() # mean loss on removed patches
return loss
def forward(self, imgs, tgts, bool_masked_pos=None, valid=None, seg_type=None, merge_between_batch=-1):
if bool_masked_pos is None:
bool_masked_pos = torch.zeros((imgs.shape[0], self.patch_embed.num_patches), dtype=torch.bool).to(imgs.device)
else:
bool_masked_pos = bool_masked_pos.flatten(1).to(torch.bool)
latent = self.forward_encoder(imgs, tgts, bool_masked_pos, seg_type, merge_between_batch=merge_between_batch)
pred = self.forward_decoder(latent) # [N, L, p*p*3]
loss = self.forward_loss(pred, tgts, bool_masked_pos, valid)
return loss, self.patchify(pred), bool_masked_pos
def seggpt_vit_large_patch16_input896x448(**kwargs):
model = SegGPT(
img_size=(896, 448), patch_size=16, embed_dim=1024, depth=24, num_heads=16,
drop_path_rate=0.1, window_size=14, qkv_bias=True,
mlp_ratio=4, norm_layer=partial(nn.LayerNorm, eps=1e-6),
window_block_indexes=(list(range(0, 2)) + list(range(3, 5)) + list(range(6, 8)) + list(range(9, 11)) + \
list(range(12, 14)), list(range(15, 17)), list(range(18, 20)), list(range(21, 23))),
residual_block_indexes=[], use_rel_pos=True, out_feature="last_feat",
decoder_embed_dim=64,
loss_func="smoothl1",
**kwargs)
return model
def get_vit_lr_decay_rate(name, lr_decay_rate=1.0, num_layers=12):
"""
Calculate lr decay rate for different ViT blocks.
Args:
name (string): parameter name.
lr_decay_rate (float): base lr decay rate.
num_layers (int): number of ViT blocks.
Returns:
lr decay rate for the given parameter.
"""
layer_id = num_layers + 1
if name.startswith("backbone"):
if ".pos_embed" in name or ".patch_embed" in name:
layer_id = 0
elif ".blocks." in name and ".residual." not in name:
layer_id = int(name[name.find(".blocks.") :].split(".")[2]) + 1
return lr_decay_rate ** (num_layers + 1 - layer_id)
+99
View File
@@ -0,0 +1,99 @@
from PIL import Image, ImageFilter, ImageEnhance, ImageOps, ImageDraw, ImageChops, ImageFont
from PIL.PngImagePlugin import PngInfo
from .models_seggpt import seggpt_vit_large_patch16_input896x448
from .seggpt_engine import inference_image_pil
import torch
import numpy as np
import math
import comfy.utils
import sys
INT = ("INT", {"default": 512,
"min": -10240,
"max": 10240,
"step": 64})
def get_image_size(IMAGE) -> tuple[int, int]:
samples = IMAGE.movedim(-1, 1)
size = samples.shape[3], samples.shape[2]
# size = size.movedim(1, -1)
return size
def convert_to_nearest_multiple_of_64(num):
return ((num + 31) // 64) * 64
import os
# 获取当前文件的目录
def prepare_model(seg_type='semantic'):
# build model
model = seggpt_vit_large_patch16_input896x448()
model.seg_type = seg_type
# load model
current_directory = os.path.dirname(os.path.abspath(__file__))
checkpoint = torch.load(os.path.join(current_directory,'seggpt_vit_large.pth'), map_location='cpu')
msg = model.load_state_dict(checkpoint['model'], strict=False)
model.eval()
return model
def pil2tensor(image):
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
class SegGPT:
upscale_methods = ["nearest-exact", "bilinear", "area"]
crop_methods = ["disabled", "center"]
def __init__(self) -> None:
pass
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"images": ("IMAGE",),
"prompt": ("IMAGE",),
"promptMask": ("IMAGE",),
}}
RETURN_TYPES = ("IMAGE","IMAGE",)
RETURN_NAMES = ("MASKS", "PREVIEW",)
FUNCTION = "doSegGPT"
CATEGORY = "SegGPT"
def doSegGPT(self, images, prompt,promptMask):
device = comfy.model_management.get_torch_device()
model = prepare_model().to(device)
prompt = Image.fromarray(np.clip(255. * prompt[0].cpu().numpy(), 0, 255).astype(np.uint8))
promptMask = Image.fromarray(np.clip(255. * promptMask[0].cpu().numpy(), 0, 255).astype(np.uint8))
results = []
resultsPrev = []
ii = 0
for image in images:
i = 255. * image.cpu().numpy()
img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8))
rImg,rImgPrev = inference_image_pil(model,device,img,[prompt],[promptMask])
rNPImg = pil2tensor(rImg)
rNPImgPrev = pil2tensor(rImgPrev)
results.append(rNPImg)
resultsPrev.append(rNPImgPrev)
ii = ii + 1
print(f"segGPT ok:{ii}")
r1 = torch.cat(results, dim=0)
r2 = torch.cat(resultsPrev, dim=0)
del model
return (r1,r2)
'''
seggpt_test = Image.open("seggpt_test.jpg")
prompt = Image.open("prompt.jpg")
promptMask = Image.open("promptMask.jpg")
device = torch.device("cuda")
model = prepare_model().to(device)
mask,preview = inference_image_pil(model,device,seggpt_test,[prompt],[promptMask])
mask.save('mask.jpg')
preview.save('preview.jpg')
#inference_image(model, device, "seggpt_test.jpg", ["prompt.jpg"], ["promptMask.jpg"], 'mask.jpg','preview.jpg')
'''
+146
View File
@@ -0,0 +1,146 @@
import os
import torch
import gradio as gr
import sys
import requests
import argparse
import torch
import torch.nn.functional as F
import numpy as np
import glob
import tqdm
import models_seggpt
from PIL import Image
imagenet_mean = np.array([0.485, 0.456, 0.406])
imagenet_std = np.array([0.229, 0.224, 0.225])
#command = f'python \
# util/painter_inference_demo2.py'
#os.system(command)
def prepare_model(chkpt_dir, arch='seggpt_vit_large_patch16_input896x448', seg_type='instance'):
# build model
model = getattr(models_seggpt, arch)()
model.seg_type = seg_type
# load model
checkpoint = torch.load(chkpt_dir, map_location='cpu')
msg = model.load_state_dict(checkpoint['model'], strict=False)
model.eval()
return model
'semantic'
device = torch.device('cuda')
model = prepare_model('seggpt_vit_large.pth', 'seggpt_vit_large_patch16_input896x448','instance' ).to(device)
imagenet_mean = np.array([0.485, 0.456, 0.406])
imagenet_std = np.array([0.229, 0.224, 0.225])
class Cache(list):
def __init__(self, max_size=0):
super().__init__()
self.max_size = max_size
def append(self, x):
if self.max_size <= 0:
return
super().append(x)
if len(self) > self.max_size:
self.pop(0)
@torch.no_grad()
def run_one_image(img, tgt, model, device):
x = torch.tensor(img)
# make it a batch-like
x = torch.einsum('nhwc->nchw', x)
tgt = torch.tensor(tgt)
# make it a batch-like
tgt = torch.einsum('nhwc->nchw', tgt)
bool_masked_pos = torch.zeros(model.patch_embed.num_patches)
bool_masked_pos[model.patch_embed.num_patches//2:] = 1
bool_masked_pos = bool_masked_pos.unsqueeze(dim=0)
valid = torch.ones_like(tgt)
if model.seg_type == 'instance':
seg_type = torch.ones([valid.shape[0], 1])
else:
seg_type = torch.zeros([valid.shape[0], 1])
feat_ensemble = 0 if len(x) > 1 else -1
_, y, mask = model(x.float().to(device), tgt.float().to(device), bool_masked_pos.to(device), valid.float().to(device), seg_type.to(device), feat_ensemble)
y = model.unpatchify(y)
y = torch.einsum('nchw->nhwc', y).detach().cpu()
output = y[0, y.shape[1]//2:, :, :]
output = torch.clip((output * imagenet_std + imagenet_mean) * 255, 0, 255)
return output
def inference_image(model, device, img_path, img2_paths, tgt2_paths):
res, hres = 448, 448
image = img_path#Image.open(img_path).convert("RGB")
input_image = np.array(image)
size = image.size
image = np.array(image.resize((res, hres))) / 255.
image_batch, target_batch = [], []
for img2_path, tgt2_path in zip(img2_paths, tgt2_paths):
img2 = img2_path#Image.open(img2_path).convert("RGB")
img2 = img2.resize((res, hres))
img2 = np.array(img2) / 255.
tgt2 = tgt2_path#Image.open(tgt2_path).convert("RGB")
tgt2 = tgt2.resize((res, hres), Image.NEAREST)
tgt2 = np.array(tgt2) / 255.
tgt = tgt2 # tgt is not available
tgt = np.concatenate((tgt2, tgt), axis=0)
img = np.concatenate((img2, image), axis=0)
assert img.shape == (2*res, res, 3), f'{img.shape}'
# normalize by ImageNet mean and std
img = img - imagenet_mean
img = img / imagenet_std
assert tgt.shape == (2*res, res, 3), f'{img.shape}'
# normalize by ImageNet mean and std
tgt = tgt - imagenet_mean
tgt = tgt / imagenet_std
image_batch.append(img)
target_batch.append(tgt)
img = np.stack(image_batch, axis=0)
tgt = np.stack(target_batch, axis=0)
"""### Run SegGPT on the image"""
# make random mask reproducible (comment out to make it change)
torch.manual_seed(2)
output = run_one_image(img, tgt, model, device)
output = F.interpolate(
output[None, ...].permute(0, 3, 1, 2),
size=[size[1], size[0]],
mode='nearest',
).permute(0, 2, 3, 1)[0].numpy()
output = Image.fromarray((input_image * (0.6 * output / 255 + 0.4)).astype(np.uint8))
return output
def predict(imagemask,image2):
img2 = imagemask["image"].convert("RGB")
tgt2 = imagemask["mask"].convert("RGB")
img = image2
return inference_image(model, device, img, [img2], [tgt2])
with gr.Blocks(css='.fixed-height.svelte-rlgzoo {height: 100%;}') as demo:
with gr.Column():
with gr.Row():
inp = [gr.ImageMask(type="pil").style(height=400),gr.Image(type="pil").style(height=400)]
out = gr.Image(type="pil")
btn = gr.Button("Run")
btn.click(fn=predict, inputs=inp, outputs=out)
demo.launch(server_name="192.168.0.100",server_port=27871)
+230
View File
@@ -0,0 +1,230 @@
import torch
import torch.nn.functional as F
import numpy as np
import cv2
from PIL import Image
imagenet_mean = np.array([0.485, 0.456, 0.406])
imagenet_std = np.array([0.229, 0.224, 0.225])
class Cache(list):
def __init__(self, max_size=0):
super().__init__()
self.max_size = max_size
def append(self, x):
if self.max_size <= 0:
return
super().append(x)
if len(self) > self.max_size:
self.pop(0)
@torch.no_grad()
def run_one_image(img, tgt, model, device):
x = torch.tensor(img)
# make it a batch-like
x = torch.einsum('nhwc->nchw', x)
tgt = torch.tensor(tgt)
# make it a batch-like
tgt = torch.einsum('nhwc->nchw', tgt)
bool_masked_pos = torch.zeros(model.patch_embed.num_patches)
bool_masked_pos[model.patch_embed.num_patches//2:] = 1
bool_masked_pos = bool_masked_pos.unsqueeze(dim=0)
valid = torch.ones_like(tgt)
if model.seg_type == 'instance':
seg_type = torch.ones([valid.shape[0], 1])
else:
seg_type = torch.zeros([valid.shape[0], 1])
feat_ensemble = 0 if len(x) > 1 else -1
_, y, mask = model(x.float().to(device), tgt.float().to(device), bool_masked_pos.to(device), valid.float().to(device), seg_type.to(device), feat_ensemble)
y = model.unpatchify(y)
y = torch.einsum('nchw->nhwc', y).detach().cpu()
output = y[0, y.shape[1]//2:, :, :]
output = torch.clip((output * imagenet_std + imagenet_mean) * 255, 0, 255)
return output
def inference_image_pil(model, device, image, prompts, promptMasks):
res, hres = 448, 448
image = image.convert("RGB")
input_image = np.array(image)
size = image.size
image = np.array(image.resize((res, hres))) / 255.
image_batch, target_batch = [], []
for prompt, promptMask in zip(prompts, promptMasks):
prompt = prompt.resize((res, hres))
prompt = np.array(prompt) / 255.
promptMask = promptMask.resize((res, hres), Image.NEAREST)
promptMask = np.array(promptMask) / 255.
tgt = np.concatenate((promptMask, promptMask), axis=0)
img = np.concatenate((prompt, image), axis=0)
assert img.shape == (2*res, res, 3), f'{img.shape}'
# normalize by ImageNet mean and std
img = img - imagenet_mean
img = img / imagenet_std
assert tgt.shape == (2*res, res, 3), f'{img.shape}'
# normalize by ImageNet mean and std
tgt = tgt - imagenet_mean
tgt = tgt / imagenet_std
image_batch.append(img)
target_batch.append(tgt)
img = np.stack(image_batch, axis=0)
tgt = np.stack(target_batch, axis=0)
"""### Run SegGPT on the image"""
# make random mask reproducible (comment out to make it change)
torch.manual_seed(2)
output = run_one_image(img, tgt, model, device)
output = F.interpolate(
output[None, ...].permute(0, 3, 1, 2),
size=[size[1], size[0]],
mode='nearest',
).permute(0, 2, 3, 1)[0].numpy()
outputx = output
output = Image.fromarray(output.astype(np.uint8))
output2 = Image.fromarray(((input_image / 2) + (outputx / 2)).astype(np.uint8))
return (output,output2)
def inference_image(model, device, img_path, img2_paths, tgt2_paths, out_path,out_path2):
res, hres = 448, 448
image = Image.open(img_path).convert("RGB")
input_image = np.array(image)
size = image.size
image = np.array(image.resize((res, hres))) / 255.
image_batch, target_batch = [], []
for img2_path, tgt2_path in zip(img2_paths, tgt2_paths):
img2 = Image.open(img2_path).convert("RGB")
img2 = img2.resize((res, hres))
img2 = np.array(img2) / 255.
tgt2 = Image.open(tgt2_path).convert("RGB")
tgt2 = tgt2.resize((res, hres), Image.NEAREST)
tgt2 = np.array(tgt2) / 255.
tgt = tgt2 # tgt is not available
tgt = np.concatenate((tgt2, tgt), axis=0)
img = np.concatenate((img2, image), axis=0)
assert img.shape == (2*res, res, 3), f'{img.shape}'
# normalize by ImageNet mean and std
img = img - imagenet_mean
img = img / imagenet_std
assert tgt.shape == (2*res, res, 3), f'{img.shape}'
# normalize by ImageNet mean and std
tgt = tgt - imagenet_mean
tgt = tgt / imagenet_std
image_batch.append(img)
target_batch.append(tgt)
img = np.stack(image_batch, axis=0)
tgt = np.stack(target_batch, axis=0)
"""### Run SegGPT on the image"""
# make random mask reproducible (comment out to make it change)
torch.manual_seed(2)
output = run_one_image(img, tgt, model, device)
output = F.interpolate(
output[None, ...].permute(0, 3, 1, 2),
size=[size[1], size[0]],
mode='nearest',
).permute(0, 2, 3, 1)[0].numpy()
outputx = output
output = Image.fromarray(output.astype(np.uint8))
output.save(out_path)
output2 = Image.fromarray(((input_image / 2) + (outputx / 2)).astype(np.uint8))
output2.save(out_path2)
def inference_video(model, device, vid_path, num_frames, img2_paths, tgt2_paths, out_path):
res, hres = 448, 448
cap = cv2.VideoCapture(vid_path)
fps = cap.get(cv2.CAP_PROP_FPS)
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
fourcc = cv2.VideoWriter_fourcc(*'mp4v')
video_writer = cv2.VideoWriter(out_path, fourcc, fps, (width, height), True)
if img2_paths is None:
_, frame = cap.read()
img2 = Image.fromarray(frame[:, :, ::-1]).convert('RGB')
else:
img2 = Image.open(img2_paths[0]).convert("RGB")
img2 = img2.resize((res, hres))
img2 = np.array(img2) / 255.
tgt2 = Image.open(tgt2_paths[0]).convert("RGB")
tgt2 = tgt2.resize((res, hres), Image.NEAREST)
tgt2 = np.array(tgt2) / 255.
frames_cache, target_cache = Cache(num_frames), Cache(num_frames)
while True:
ret, frame = cap.read()
if not ret:
break
image_batch, target_batch = [], []
image = Image.fromarray(frame[:, :, ::-1]).convert('RGB')
input_image = np.array(image)
size = image.size
image = np.array(image.resize((res, hres))) / 255.
for prompt, target in zip([img2] + frames_cache, [tgt2] + target_cache):
tgt = target # tgt is not available
tgt = np.concatenate((target, tgt), axis=0)
img = np.concatenate((prompt, image), axis=0)
assert img.shape == (2*res, res, 3), f'{img.shape}'
# normalize by ImageNet mean and std
img = img - imagenet_mean
img = img / imagenet_std
assert tgt.shape == (2*res, res, 3), f'{img.shape}'
# normalize by ImageNet mean and std
tgt = tgt - imagenet_mean
tgt = tgt / imagenet_std
image_batch.append(img)
target_batch.append(tgt)
img = np.stack(image_batch, axis=0)
tgt = np.stack(target_batch, axis=0)
torch.manual_seed(2)
output = run_one_image(img, tgt, model, device)
frames_cache.append(image)
target_cache.append(
output.mean(-1) \
.gt(128).float() \
.unsqueeze(-1).expand(-1, -1, 3) \
.numpy()
)
output = F.interpolate(
output[None, ...].permute(0, 3, 1, 2),
size=[size[1], size[0]],
mode='nearest',
).permute(0, 2, 3, 1)[0].numpy()
output = input_image * (0.6 * output / 255 + 0.4)
video_writer.write(np.ascontiguousarray(output.astype(np.uint8)[:, :, ::-1]))
video_writer.release()
+74
View File
@@ -0,0 +1,74 @@
import os
import argparse
import torch
import numpy as np
from seggpt_engine import inference_image, inference_video
import models_seggpt
imagenet_mean = np.array([0.485, 0.456, 0.406])
imagenet_std = np.array([0.229, 0.224, 0.225])
def get_args_parser():
parser = argparse.ArgumentParser('SegGPT inference', add_help=False)
parser.add_argument('--ckpt_path', type=str, help='path to ckpt',
default='seggpt_vit_large.pth')
parser.add_argument('--model', type=str, help='dir to ckpt',
default='seggpt_vit_large_patch16_input896x448')
parser.add_argument('--input_image', type=str, help='path to input image to be tested',
default=None)
parser.add_argument('--input_video', type=str, help='path to input video to be tested',
default=None)
parser.add_argument('--num_frames', type=int, help='number of prompt frames in video',
default=0)
parser.add_argument('--prompt_image', type=str, nargs='+', help='path to prompt image',
default=None)
parser.add_argument('--prompt_target', type=str, nargs='+', help='path to prompt target',
default=None)
parser.add_argument('--seg_type', type=str, help='embedding for segmentation types',
choices=['instance', 'semantic'], default='instance')
parser.add_argument('--device', type=str, help='cuda or cpu',
default='cuda')
parser.add_argument('--output_dir', type=str, help='path to output',
default='./')
return parser.parse_args()
def prepare_model(chkpt_dir, arch='seggpt_vit_large_patch16_input896x448', seg_type='instance'):
# build model
model = getattr(models_seggpt, arch)()
model.seg_type = seg_type
# load model
checkpoint = torch.load(chkpt_dir, map_location='cpu')
msg = model.load_state_dict(checkpoint['model'], strict=False)
model.eval()
return model
if __name__ == '__main__':
args = get_args_parser()
device = torch.device(args.device)
model = prepare_model(args.ckpt_path, args.model, args.seg_type).to(device)
print('Model loaded.')
assert args.input_image or args.input_video and not (args.input_image and args.input_video)
if args.input_image is not None:
assert args.prompt_image is not None and args.prompt_target is not None
img_name = os.path.basename(args.input_image)
out_path = os.path.join(args.output_dir, "output_" + '.'.join(img_name.split('.')[:-1]) + '.png')
inference_image(model, device, args.input_image, args.prompt_image, args.prompt_target, out_path)
if args.input_video is not None:
assert args.prompt_target is not None and len(args.prompt_target) == 1
vid_name = os.path.basename(args.input_video)
out_path = os.path.join(args.output_dir, "output_" + '.'.join(vid_name.split('.')[:-1]) + '.mp4')
inference_video(model, device, args.input_video, args.num_frames, args.prompt_image, args.prompt_target, out_path)
print('Finished.')
+85
View File
@@ -0,0 +1,85 @@
import os
import argparse
import torch
import numpy as np
from seggpt_engine import inference_image, inference_video
import models_seggpt
imagenet_mean = np.array([0.485, 0.456, 0.406])
imagenet_std = np.array([0.229, 0.224, 0.225])
def get_args_parser():
parser = argparse.ArgumentParser('SegGPT inference', add_help=False)
parser.add_argument('--ckpt_path', type=str, help='path to ckpt',
default='seggpt_vit_large.pth')
parser.add_argument('--model', type=str, help='dir to ckpt',
default='seggpt_vit_large_patch16_input896x448')
parser.add_argument('--input_image', type=str, help='path to input image to be tested',
default=None)
parser.add_argument('--input_image_path', type=str, help='path to input path to be tested',
default=None)
parser.add_argument('--input_video', type=str, help='path to input video to be tested',
default=None)
parser.add_argument('--num_frames', type=int, help='number of prompt frames in video',
default=0)
parser.add_argument('--prompt_image', type=str, nargs='+', help='path to prompt image',
default=None)
parser.add_argument('--prompt_target', type=str, nargs='+', help='path to prompt target',
default=None)
parser.add_argument('--seg_type', type=str, help='embedding for segmentation types',
choices=['instance', 'semantic'], default='semantic')
parser.add_argument('--device', type=str, help='cuda or cpu',
default='cuda')
parser.add_argument('--output_dir', type=str, help='path to output',
default='./')
return parser.parse_args()
def prepare_model(chkpt_dir, arch='seggpt_vit_large_patch16_input896x448', seg_type='semantic'):
# build model
model = getattr(models_seggpt, arch)()
model.seg_type = seg_type
# load model
checkpoint = torch.load(chkpt_dir, map_location='cpu')
msg = model.load_state_dict(checkpoint['model'], strict=False)
model.eval()
return model
if __name__ == '__main__':
args = get_args_parser()
device = torch.device(args.device)
model = prepare_model(args.ckpt_path, args.model, args.seg_type).to(device)
print('Model loaded.')
if args.input_image_path is not None:
for file_name in os.listdir(args.input_image_path):
f = os.path.join(args.input_image_path,file_name)
print(f)
img_name = os.path.basename(f)
out_path = os.path.join(args.output_dir, file_name)
out_path2 = os.path.join(args.output_dir + '_', file_name)
inference_image(model, device,f, args.prompt_image,args.prompt_target, out_path, out_path2)
exit(0)
assert args.input_image or args.input_video and not (args.input_image and args.input_video)
if args.input_image is not None:
assert args.prompt_image is not None
assert args.prompt_target is not None
img_name = os.path.basename(args.input_image)
out_path = os.path.join(args.output_dir, "output_" + '.'.join(img_name.split('.')[:-1]) + '.png')
inference_image(model, device, args.input_image, args.prompt_image, args.prompt_target, out_path)
if args.input_video is not None:
assert args.prompt_target is not None and len(args.prompt_target) == 1
vid_name = os.path.basename(args.input_video)
out_path = os.path.join(args.output_dir, "output_" + '.'.join(vid_name.split('.')[:-1]) + '.mp4')
inference_video(model, device, args.input_video, args.num_frames, args.prompt_image, args.prompt_target, out_path)
print('Finished.')
+75
View File
@@ -0,0 +1,75 @@
from PIL import Image, ImageFilter, ImageEnhance, ImageOps, ImageDraw, ImageChops, ImageFont
from PIL.PngImagePlugin import PngInfo
from models_seggpt import seggpt_vit_large_patch16_input896x448
from seggpt_engine import inference_image_pil,inference_image
import torch
import numpy as np
import math
import sys
INT = ("INT", {"default": 512,
"min": -10240,
"max": 10240,
"step": 64})
def get_image_size(IMAGE) -> tuple[int, int]:
samples = IMAGE.movedim(-1, 1)
size = samples.shape[3], samples.shape[2]
# size = size.movedim(1, -1)
return size
def convert_to_nearest_multiple_of_64(num):
return ((num + 31) // 64) * 64
import os
# 获取当前文件的目录
def prepare_model(seg_type='semantic'):
# build model
model = seggpt_vit_large_patch16_input896x448()
model.seg_type = seg_type
# load model
current_directory = os.path.dirname(os.path.abspath(__file__))
checkpoint = torch.load(os.path.join(current_directory,'seggpt_vit_large.pth'), map_location='cpu')
msg = model.load_state_dict(checkpoint['model'], strict=False)
model.eval()
return model
class SegGPT:
upscale_methods = ["nearest-exact", "bilinear", "area"]
crop_methods = ["disabled", "center"]
def __init__(self) -> None:
pass
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"images": ("IMAGE",),
"prompt": ("IMAGE",),
"promptMask": ("IMAGE",),
}}
RETURN_TYPES = ("IMAGE","IMAGE",)
RETURN_NAMES = ("MASKS", "PREVIEW",)
FUNCTION = "doSegGPT"
CATEGORY = "SegGPT"
def doSegGPT(self, images, prompt,promptMask):
model = prepare_model().to(device)
prompt = Image.fromarray(np.clip(255. * prompt[0].cpu().numpy(), 0, 255).astype(np.uint8))
promptMask = Image.fromarray(np.clip(255. * promptMask[0].cpu().numpy(), 0, 255).astype(np.uint8))
results = []
resultsPrev = []
for image in images:
i = 255. * image.cpu().numpy()
img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8))
rNPImg,rNPImgPrev = np.array(inference_image_pil(model,comfy.model_management.get_torch_device(),img,[prompt],[promptMask])) / 255.
results.append(rNPImg)
resultsPrev.append(rNPImgPrev)
r1 = np.array(results)
r1 = np.array(resultsPrev)
return (r1,r2)
+19
View File
@@ -0,0 +1,19 @@
import os
f = r'C:\telegramDownload\oldWork\定制\jc\jc_02'
if not os.path.exists(f'{f}/preview'):
os.makedirs(f'{f}/preview')
if not os.path.exists(f'{f}/preview_'):
os.makedirs(f'{f}/preview_')
command = f'python seggpt_inference_batch.py \
--input_image_path {f}/frames \
--prompt_image {f}/prompt.jpg {f}/prompt2.jpg {f}/prompt3.jpg \
--prompt_target {f}/promptMask.jpg {f}/promptMask2.jpg {f}/promptMask3.jpg \
--output_dir {f}/preview'
command = f'python seggpt_inference_batch.py \
--input_image_path {f}/frames2 \
--prompt_image {f}/prompt.jpg \
--prompt_target {f}/promptMask.jpg \
--output_dir {f}/preview'
os.system(command)
+209
View File
@@ -0,0 +1,209 @@
# Copyright (c) Facebook, Inc. and its affiliates. All Rights Reserved
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
__all__ = [
"window_partition",
"window_unpartition",
"add_decomposed_rel_pos",
"get_abs_pos",
"PatchEmbed",
]
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
def get_rel_pos(q_size, k_size, rel_pos):
"""
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, q, rel_pos_h, rel_pos_w, q_size, k_size):
"""
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
def get_abs_pos(abs_pos, has_cls_token, hw):
"""
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, w = hw
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)
class PatchEmbed(nn.Module):
"""
Image to Patch Embedding.
"""
def __init__(
self, kernel_size=(16, 16), stride=(16, 16), padding=(0, 0), in_chans=3, embed_dim=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):
x = self.proj(x)
# B C H W -> B H W C
x = x.permute(0, 2, 3, 1)
return x
class LayerNorm2D(nn.Module):
"""
A LayerNorm variant, popularized by Transformers, that performs point-wise mean and
variance normalization over the channel dimension for inputs that have shape
(batch_size, channels, height, width).
https://github.com/facebookresearch/ConvNeXt/blob/d1fa8f6fef0a165b27399986cc2bdacc92777e40/models/convnext.py#L119 # noqa B950
"""
def __init__(self, normalized_shape, eps=1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(normalized_shape))
self.bias = nn.Parameter(torch.zeros(normalized_shape))
self.eps = eps
self.normalized_shape = (normalized_shape,)
def forward(self, x):
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