Add files via upload
This commit is contained in:
@@ -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
@@ -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']
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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.')
|
||||
@@ -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.')
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user