diff --git a/README.md b/README.md new file mode 100644 index 0000000..dc349e9 --- /dev/null +++ b/README.md @@ -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 ./ +``` + + diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..dc08efb --- /dev/null +++ b/__init__.py @@ -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'] diff --git a/models_seggpt.py b/models_seggpt.py new file mode 100644 index 0000000..6bb8a04 --- /dev/null +++ b/models_seggpt.py @@ -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) + diff --git a/seggpt.py b/seggpt.py new file mode 100644 index 0000000..7f28180 --- /dev/null +++ b/seggpt.py @@ -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') +''' + diff --git a/seggpt_app.py b/seggpt_app.py new file mode 100644 index 0000000..7039e68 --- /dev/null +++ b/seggpt_app.py @@ -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) diff --git a/seggpt_engine.py b/seggpt_engine.py new file mode 100644 index 0000000..ebf6aa6 --- /dev/null +++ b/seggpt_engine.py @@ -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() diff --git a/seggpt_inference.py b/seggpt_inference.py new file mode 100644 index 0000000..bc674a9 --- /dev/null +++ b/seggpt_inference.py @@ -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.') diff --git a/seggpt_inference_batch.py b/seggpt_inference_batch.py new file mode 100644 index 0000000..ec9e1c2 --- /dev/null +++ b/seggpt_inference_batch.py @@ -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.') diff --git a/seggpt_test.py b/seggpt_test.py new file mode 100644 index 0000000..efbdcaf --- /dev/null +++ b/seggpt_test.py @@ -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) + diff --git a/test.py b/test.py new file mode 100644 index 0000000..f4fbb5f --- /dev/null +++ b/test.py @@ -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) diff --git a/util/vitdet_utils.py b/util/vitdet_utils.py new file mode 100644 index 0000000..37b1742 --- /dev/null +++ b/util/vitdet_utils.py @@ -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 \ No newline at end of file