update v1.1.0
This commit is contained in:
@@ -3,9 +3,21 @@
|
||||
from scepter.modules.annotator.base_annotator import GeneralAnnotator
|
||||
from scepter.modules.annotator.canny import CannyAnnotator
|
||||
from scepter.modules.annotator.color import ColorAnnotator
|
||||
from scepter.modules.annotator.degradation import DegradationAnnotator
|
||||
from scepter.modules.annotator.doodle import DoodleAnnotator
|
||||
from scepter.modules.annotator.gray import GrayAnnotator
|
||||
from scepter.modules.annotator.hed import HedAnnotator
|
||||
from scepter.modules.annotator.identity import IdentityAnnotator
|
||||
from scepter.modules.annotator.informative_drawing import (
|
||||
InfoDrawAnimeAnnotator, InfoDrawContourAnnotator,
|
||||
InfoDrawOpenSketchAnnotator)
|
||||
from scepter.modules.annotator.inpainting import InpaintingAnnotator
|
||||
from scepter.modules.annotator.invert import InvertAnnotator
|
||||
from scepter.modules.annotator.midas_op import MidasDetector
|
||||
from scepter.modules.annotator.mlsd_op import MLSDdetector
|
||||
from scepter.modules.annotator.openpose import OpenposeAnnotator
|
||||
from scepter.modules.annotator.outpainting import OutpaintingAnnotator, OutpaintingResize
|
||||
from scepter.modules.annotator.pidinet import PiDiAnnotator
|
||||
from scepter.modules.annotator.segmentation import ESAMAnnotator
|
||||
from scepter.modules.annotator.sketch import SketchAnnotator
|
||||
from scepter.modules.annotator.lama import LamaAnnotator
|
||||
|
||||
@@ -0,0 +1,135 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import math
|
||||
import random
|
||||
from abc import ABCMeta
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from PIL import Image
|
||||
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||
from scepter.modules.annotator.registry import ANNOTATORS
|
||||
from scepter.modules.utils.config import Config, dict_to_yaml
|
||||
|
||||
|
||||
def gaussian_noise_op(im, v):
|
||||
from basicsr.data.degradations import random_add_gaussian_noise
|
||||
noise_level = v.get('noise_level', [10, 20])
|
||||
out = random_add_gaussian_noise(
|
||||
im,
|
||||
sigma_range=noise_level,
|
||||
clip=True,
|
||||
rounds=False,
|
||||
gray_prob=0.4,
|
||||
)
|
||||
out = np.clip(out, 0.0, 1.0)
|
||||
return out
|
||||
|
||||
|
||||
def resize_op(im, v):
|
||||
scale = v.get('scale', [0.5, 0.8])
|
||||
h, w = im.shape[:2]
|
||||
scale = random.uniform(scale[0], scale[1])
|
||||
h_, w_ = int(h * scale), int(w * scale)
|
||||
mode = v.get('mode', 'nearest')
|
||||
if mode == 'nearest':
|
||||
interpolation = cv2.INTER_NEAREST
|
||||
elif mode == 'bilinear':
|
||||
interpolation = cv2.INTER_LINEAR
|
||||
elif mode == 'bicubic':
|
||||
interpolation = cv2.INTER_CUBIC
|
||||
else:
|
||||
interpolation = cv2.INTER_NEAREST
|
||||
im = cv2.resize(im, (w_, h_), interpolation=interpolation)
|
||||
out = cv2.resize(im, (w, h), interpolation=interpolation)
|
||||
out = np.clip(out, 0.0, 1.0)
|
||||
return out
|
||||
|
||||
|
||||
def jpeg_op(im, v):
|
||||
from basicsr.data.degradations import add_jpg_compression
|
||||
jpeg_level = v.get('jpeg_level', [50, 75])
|
||||
v = int(random.uniform(jpeg_level[0], jpeg_level[1]))
|
||||
out = add_jpg_compression(im, v)
|
||||
out = np.clip(out, 0.0, 1.0)
|
||||
return out
|
||||
|
||||
|
||||
def gaussian_blur_op(im, v):
|
||||
from basicsr.data.degradations import random_mixed_kernels
|
||||
kernel_range = v.get('kernel_size', [7, 9])
|
||||
kernel_size = random.choice(kernel_range)
|
||||
kernel_size = min(int(kernel_size) // 2 * 2 + 1, 21)
|
||||
blur_sigma = v.get('sigma', [0.9, 1.0])
|
||||
kernel = random_mixed_kernels(
|
||||
('iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso',
|
||||
'plateau_aniso'), (0.45, 0.25, 0.12, 0.03, 0.12, 0.03),
|
||||
kernel_size,
|
||||
blur_sigma,
|
||||
blur_sigma, [-math.pi, math.pi], [0.5, 2.0], [1, 1.5],
|
||||
noise_range=None)
|
||||
|
||||
pad_size = (21 - kernel_size) // 2
|
||||
kernel = np.pad(kernel, ((pad_size, pad_size), (pad_size, pad_size)))
|
||||
out = cv2.filter2D(im, -1, kernel)
|
||||
out = np.clip(out, 0.0, 1.0)
|
||||
return out
|
||||
|
||||
|
||||
@ANNOTATORS.register_class()
|
||||
class DegradationAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
para_dict = {}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.params = cfg.get('PARAMS', {
|
||||
'gaussian_noise': {},
|
||||
'resize': {},
|
||||
'jpeg': {},
|
||||
'gaussian_blur': {},
|
||||
})
|
||||
if not isinstance(self.params, dict):
|
||||
self.params = Config.get_dict(self.params)
|
||||
self.random_degradation = cfg.get('RANDOM_DEGRADATION', False)
|
||||
|
||||
def forward(self, image):
|
||||
if isinstance(image, Image.Image):
|
||||
image = np.array(image)
|
||||
elif isinstance(image, torch.Tensor):
|
||||
image = image.detach().cpu().numpy()
|
||||
elif isinstance(image, np.ndarray):
|
||||
image = image.copy()
|
||||
else:
|
||||
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
|
||||
if np.max(image) > 1.0:
|
||||
image = (image / 255.).astype(np.float32)
|
||||
|
||||
degradation_list = list(self.params.keys())
|
||||
if self.random_degradation:
|
||||
random.shuffle(degradation_list)
|
||||
|
||||
for degradation_type in degradation_list:
|
||||
if degradation_type == 'gaussian_noise':
|
||||
image = gaussian_noise_op(image, self.params[degradation_type])
|
||||
elif degradation_type == 'resize':
|
||||
image = resize_op(image, self.params[degradation_type])
|
||||
elif degradation_type == 'jpeg':
|
||||
image = jpeg_op(image, self.params[degradation_type])
|
||||
elif degradation_type == 'gaussian_blur':
|
||||
image = gaussian_blur_op(image, self.params[degradation_type])
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f'ERROR: degradation_type: {degradation_type} is invalid.')
|
||||
image = (image * 255.0).astype(np.uint8)
|
||||
|
||||
assert len(image.shape) < 4
|
||||
return image
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('ANNOTATORS',
|
||||
__class__.__name__,
|
||||
DegradationAnnotator.para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,51 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
import math
|
||||
from abc import ABCMeta
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torchvision.transforms as TT
|
||||
from einops import rearrange
|
||||
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||
from scepter.modules.annotator.registry import ANNOTATORS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
|
||||
@ANNOTATORS.register_class()
|
||||
class DoodleAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
para_dict = {}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.processor_type = cfg.get('PROCESSOR_TYPE', 'pidinet_sketch')
|
||||
processor_cfg = cfg.get('PROCESSOR_CFG', None)
|
||||
if self.processor_type == 'pidinet_sketch':
|
||||
self.pidinet_ins = ANNOTATORS.build(processor_cfg[0])
|
||||
self.sketch_ins = ANNOTATORS.build(processor_cfg[1])
|
||||
else:
|
||||
raise 'Unsurpport PROCESSOR for DoodleAnnotator'
|
||||
|
||||
@torch.no_grad()
|
||||
@torch.inference_mode()
|
||||
@torch.autocast('cuda', enabled=False)
|
||||
def forward(self, image):
|
||||
if self.processor_type == 'pidinet_sketch':
|
||||
pidinet_res = self.pidinet_ins(image)
|
||||
sketch_res = self.sketch_ins(pidinet_res)
|
||||
doodle_res = sketch_res
|
||||
else:
|
||||
raise 'Unsurpport PROCESSOR for DoodleAnnotator'
|
||||
return doodle_res
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('ANNOTATORS',
|
||||
__class__.__name__,
|
||||
DoodleAnnotator.para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,38 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from abc import ABCMeta
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||
from scepter.modules.annotator.registry import ANNOTATORS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
|
||||
|
||||
@ANNOTATORS.register_class()
|
||||
class GrayAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
para_dict = {}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
|
||||
def forward(self, image):
|
||||
if isinstance(image, Image.Image):
|
||||
image = np.array(image)
|
||||
elif isinstance(image, torch.Tensor):
|
||||
image = image.detach().cpu().numpy()
|
||||
elif isinstance(image, np.ndarray):
|
||||
image = image.copy()
|
||||
else:
|
||||
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
|
||||
gray_map = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)
|
||||
return gray_map[..., None].repeat(3, axis=2)
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('ANNOTATORS',
|
||||
__class__.__name__,
|
||||
GrayAnnotator.para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,178 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from abc import ABCMeta
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torchvision
|
||||
import torchvision.transforms as transforms
|
||||
from einops import rearrange
|
||||
from PIL import Image
|
||||
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||
from scepter.modules.annotator.registry import ANNOTATORS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from torchvision.transforms import InterpolationMode
|
||||
|
||||
norm_layer = nn.InstanceNorm2d
|
||||
|
||||
|
||||
class ResidualBlock(nn.Module):
|
||||
def __init__(self, in_features):
|
||||
super(ResidualBlock, self).__init__()
|
||||
|
||||
conv_block = [
|
||||
nn.ReflectionPad2d(1),
|
||||
nn.Conv2d(in_features, in_features, 3),
|
||||
norm_layer(in_features),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.ReflectionPad2d(1),
|
||||
nn.Conv2d(in_features, in_features, 3),
|
||||
norm_layer(in_features)
|
||||
]
|
||||
|
||||
self.conv_block = nn.Sequential(*conv_block)
|
||||
|
||||
def forward(self, x):
|
||||
return x + self.conv_block(x)
|
||||
|
||||
|
||||
class ContourInference(nn.Module):
|
||||
def __init__(self, input_nc, output_nc, n_residual_blocks=9, sigmoid=True):
|
||||
super(ContourInference, self).__init__()
|
||||
|
||||
# Initial convolution block
|
||||
model0 = [
|
||||
nn.ReflectionPad2d(3),
|
||||
nn.Conv2d(input_nc, 64, 7),
|
||||
norm_layer(64),
|
||||
nn.ReLU(inplace=True)
|
||||
]
|
||||
self.model0 = nn.Sequential(*model0)
|
||||
|
||||
# Downsampling
|
||||
model1 = []
|
||||
in_features = 64
|
||||
out_features = in_features * 2
|
||||
for _ in range(2):
|
||||
model1 += [
|
||||
nn.Conv2d(in_features, out_features, 3, stride=2, padding=1),
|
||||
norm_layer(out_features),
|
||||
nn.ReLU(inplace=True)
|
||||
]
|
||||
in_features = out_features
|
||||
out_features = in_features * 2
|
||||
self.model1 = nn.Sequential(*model1)
|
||||
|
||||
model2 = []
|
||||
# Residual blocks
|
||||
for _ in range(n_residual_blocks):
|
||||
model2 += [ResidualBlock(in_features)]
|
||||
self.model2 = nn.Sequential(*model2)
|
||||
|
||||
# Upsampling
|
||||
model3 = []
|
||||
out_features = in_features // 2
|
||||
for _ in range(2):
|
||||
model3 += [
|
||||
nn.ConvTranspose2d(in_features,
|
||||
out_features,
|
||||
3,
|
||||
stride=2,
|
||||
padding=1,
|
||||
output_padding=1),
|
||||
norm_layer(out_features),
|
||||
nn.ReLU(inplace=True)
|
||||
]
|
||||
in_features = out_features
|
||||
out_features = in_features // 2
|
||||
self.model3 = nn.Sequential(*model3)
|
||||
|
||||
# Output layer
|
||||
model4 = [nn.ReflectionPad2d(3), nn.Conv2d(64, output_nc, 7)]
|
||||
if sigmoid:
|
||||
model4 += [nn.Sigmoid()]
|
||||
|
||||
self.model4 = nn.Sequential(*model4)
|
||||
|
||||
def forward(self, x, cond=None):
|
||||
out = self.model0(x)
|
||||
out = self.model1(out)
|
||||
out = self.model2(out)
|
||||
out = self.model3(out)
|
||||
out = self.model4(out)
|
||||
|
||||
return out
|
||||
|
||||
|
||||
@ANNOTATORS.register_class()
|
||||
class InfoDrawContourAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
para_dict = {}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
input_nc = cfg.get('INPUT_NC', 3)
|
||||
output_nc = cfg.get('OUTPUT_NC', 1)
|
||||
n_residual_blocks = cfg.get('N_RESIDUAL_BLOCKS', 3)
|
||||
sigmoid = cfg.get('SIGMOID', True)
|
||||
pretrained_model = cfg.get('PRETRAINED_MODEL', None)
|
||||
|
||||
self.model = ContourInference(input_nc, output_nc, n_residual_blocks,
|
||||
sigmoid)
|
||||
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
|
||||
self.model.load_state_dict(torch.load(local_path))
|
||||
self.model = self.model.eval().requires_grad_(False).to(we.device_id)
|
||||
|
||||
@torch.no_grad()
|
||||
@torch.inference_mode()
|
||||
@torch.autocast('cuda', enabled=False)
|
||||
def forward(self, image):
|
||||
is_batch = False if len(image.shape) == 3 else True
|
||||
if isinstance(image, torch.Tensor):
|
||||
if len(image.shape) == 3:
|
||||
image = rearrange(image, 'h w c -> 1 c h w')
|
||||
B, C, H, W = image.shape
|
||||
elif len(image.shape) == 4:
|
||||
B, C, H, W = image.shape
|
||||
else:
|
||||
raise "Unsurpport input image's shape"
|
||||
elif isinstance(image, np.ndarray):
|
||||
image = torch.from_numpy(image.copy()).float()
|
||||
if len(image.shape) == 3:
|
||||
image = rearrange(image, 'h w c -> 1 c h w')
|
||||
B, C, H, W = image.shape
|
||||
elif len(image.shape) == 4:
|
||||
B, C, H, W = image.shape
|
||||
else:
|
||||
raise "Unsurpport input image's shape"
|
||||
else:
|
||||
raise "Unsurpport input image's type"
|
||||
|
||||
image = image.float().div(255).to(we.device_id)
|
||||
contour_map = self.model(image)
|
||||
contour_map = (contour_map.squeeze(dim=1) * 255.0).clip(
|
||||
0, 255).cpu().numpy().astype(np.uint8)
|
||||
contour_map = contour_map[..., None].repeat(3, -1)
|
||||
if not is_batch:
|
||||
contour_map = contour_map.squeeze()
|
||||
return contour_map
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('ANNOTATORS',
|
||||
__class__.__name__,
|
||||
InfoDrawContourAnnotator.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@ANNOTATORS.register_class()
|
||||
class InfoDrawAnimeAnnotator(InfoDrawContourAnnotator):
|
||||
pass
|
||||
|
||||
|
||||
@ANNOTATORS.register_class()
|
||||
class InfoDrawOpenSketchAnnotator(InfoDrawContourAnnotator):
|
||||
pass
|
||||
@@ -0,0 +1,271 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import math
|
||||
import random
|
||||
from abc import ABCMeta
|
||||
from enum import Enum
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image, ImageDraw
|
||||
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||
from scepter.modules.annotator.registry import ANNOTATORS
|
||||
from scepter.modules.utils.config import Config, dict_to_yaml
|
||||
|
||||
|
||||
def invert_image(im):
|
||||
im_arr = np.array(im)
|
||||
mask_1 = im_arr == 0
|
||||
mask_2 = im_arr == 255
|
||||
im_arr[mask_1] = 255
|
||||
im_arr[mask_2] = 0
|
||||
new_im = im_arr
|
||||
return new_im
|
||||
|
||||
|
||||
class DrawMethod(Enum):
|
||||
LINE = 'line'
|
||||
CIRCLE = 'circle'
|
||||
SQUARE = 'square'
|
||||
|
||||
|
||||
def make_random_irregular_mask(shape,
|
||||
max_angle=4,
|
||||
max_len=60,
|
||||
max_width=20,
|
||||
min_times=0,
|
||||
max_times=10,
|
||||
draw_method=DrawMethod.LINE):
|
||||
draw_method = DrawMethod(draw_method)
|
||||
|
||||
height, width = shape
|
||||
mask = np.zeros((height, width), np.float32)
|
||||
times = np.random.randint(min_times, max_times + 1)
|
||||
for i in range(times):
|
||||
start_x = np.random.randint(width)
|
||||
start_y = np.random.randint(height)
|
||||
for j in range(1 + np.random.randint(5)):
|
||||
angle = 0.01 + np.random.randint(max_angle)
|
||||
if i % 2 == 0:
|
||||
angle = 2 * 3.1415926 - angle
|
||||
length = 10 + np.random.randint(max_len)
|
||||
brush_w = 5 + np.random.randint(max_width)
|
||||
end_x = np.clip(
|
||||
(start_x + length * np.sin(angle)).astype(np.int32), 0, width)
|
||||
end_y = np.clip(
|
||||
(start_y + length * np.cos(angle)).astype(np.int32), 0, height)
|
||||
if draw_method == DrawMethod.LINE:
|
||||
cv2.line(mask, (start_x, start_y), (end_x, end_y), 1.0,
|
||||
brush_w)
|
||||
elif draw_method == DrawMethod.CIRCLE:
|
||||
cv2.circle(mask, (start_x, start_y),
|
||||
radius=brush_w,
|
||||
color=1.,
|
||||
thickness=-1)
|
||||
elif draw_method == DrawMethod.SQUARE:
|
||||
radius = brush_w // 2
|
||||
mask[start_y - radius:start_y + radius,
|
||||
start_x - radius:start_x + radius] = 1
|
||||
start_x, start_y = end_x, end_y
|
||||
return mask[None, ...]
|
||||
|
||||
|
||||
class RandomIrregularMaskGenerator:
|
||||
def __init__(self,
|
||||
max_angle=4,
|
||||
max_len=60,
|
||||
max_width=20,
|
||||
min_times=0,
|
||||
max_times=10,
|
||||
ramp_kwargs=None,
|
||||
draw_method=DrawMethod.LINE):
|
||||
self.max_angle = max_angle
|
||||
self.max_len = max_len
|
||||
self.max_width = max_width
|
||||
self.min_times = min_times
|
||||
self.max_times = max_times
|
||||
self.draw_method = draw_method
|
||||
self.ramp = None
|
||||
# self.ramp = LinearRamp(**ramp_kwargs) if ramp_kwargs is not None else None
|
||||
|
||||
def __call__(self, img, iter_i=None, raw_image=None):
|
||||
coef = self.ramp(iter_i) if (self.ramp is not None) and (
|
||||
iter_i is not None) else 1
|
||||
cur_max_len = int(max(1, self.max_len * coef))
|
||||
cur_max_width = int(max(1, self.max_width * coef))
|
||||
cur_max_times = int(self.min_times + 1 +
|
||||
(self.max_times - self.min_times) * coef)
|
||||
return make_random_irregular_mask(img.shape[1:],
|
||||
max_angle=self.max_angle,
|
||||
max_len=cur_max_len,
|
||||
max_width=cur_max_width,
|
||||
min_times=self.min_times,
|
||||
max_times=cur_max_times,
|
||||
draw_method=self.draw_method)
|
||||
|
||||
|
||||
def make_random_rectangle_mask(shape,
|
||||
margin=10,
|
||||
bbox_min_size=30,
|
||||
bbox_max_size=100,
|
||||
min_times=0,
|
||||
max_times=3):
|
||||
height, width = shape
|
||||
mask = np.zeros((height, width), np.float32)
|
||||
bbox_max_size = min(bbox_max_size, height - margin * 2, width - margin * 2)
|
||||
times = np.random.randint(min_times, max_times + 1)
|
||||
for i in range(times):
|
||||
box_width = np.random.randint(bbox_min_size, bbox_max_size)
|
||||
box_height = np.random.randint(bbox_min_size, bbox_max_size)
|
||||
start_x = np.random.randint(margin, width - margin - box_width + 1)
|
||||
start_y = np.random.randint(margin, height - margin - box_height + 1)
|
||||
mask[start_y:start_y + box_height, start_x:start_x + box_width] = 1
|
||||
return mask[None, ...]
|
||||
|
||||
|
||||
class RandomRectangleMaskGenerator:
|
||||
def __init__(self,
|
||||
margin=10,
|
||||
bbox_min_size=30,
|
||||
bbox_max_size=100,
|
||||
min_times=0,
|
||||
max_times=3,
|
||||
ramp_kwargs=None):
|
||||
self.margin = margin
|
||||
self.bbox_min_size = bbox_min_size
|
||||
self.bbox_max_size = bbox_max_size
|
||||
self.min_times = min_times
|
||||
self.max_times = max_times
|
||||
self.ramp = None
|
||||
# self.ramp = LinearRamp(**ramp_kwargs) if ramp_kwargs is not None else None
|
||||
|
||||
def __call__(self, img, iter_i=None, raw_image=None):
|
||||
coef = self.ramp(iter_i) if (self.ramp is not None) and (
|
||||
iter_i is not None) else 1
|
||||
cur_bbox_max_size = int(self.bbox_min_size + 1 +
|
||||
(self.bbox_max_size - self.bbox_min_size) *
|
||||
coef)
|
||||
cur_max_times = int(self.min_times +
|
||||
(self.max_times - self.min_times) * coef)
|
||||
return make_random_rectangle_mask(img.shape[1:],
|
||||
margin=self.margin,
|
||||
bbox_min_size=self.bbox_min_size,
|
||||
bbox_max_size=cur_bbox_max_size,
|
||||
min_times=self.min_times,
|
||||
max_times=cur_max_times)
|
||||
|
||||
|
||||
class MixedMaskGenerator:
|
||||
def __init__(self,
|
||||
irregular_proba=1 / 3,
|
||||
irregular_kwargs=None,
|
||||
box_proba=1 / 3,
|
||||
box_kwargs=None,
|
||||
invert_proba=0):
|
||||
self.probas = []
|
||||
self.gens = []
|
||||
|
||||
if irregular_proba > 0:
|
||||
self.probas.append(irregular_proba)
|
||||
if irregular_kwargs is None:
|
||||
irregular_kwargs = {}
|
||||
else:
|
||||
irregular_kwargs = dict(irregular_kwargs)
|
||||
irregular_kwargs['draw_method'] = DrawMethod.LINE
|
||||
self.gens.append(RandomIrregularMaskGenerator(**irregular_kwargs))
|
||||
|
||||
if box_proba > 0:
|
||||
self.probas.append(box_proba)
|
||||
if box_kwargs is None:
|
||||
box_kwargs = {}
|
||||
self.gens.append(RandomRectangleMaskGenerator(**box_kwargs))
|
||||
|
||||
self.probas = np.array(self.probas, dtype='float32')
|
||||
self.probas /= self.probas.sum()
|
||||
self.invert_proba = invert_proba
|
||||
|
||||
def __call__(self, img, iter_i=None, raw_image=None):
|
||||
kind = np.random.choice(len(self.probas), p=self.probas)
|
||||
gen = self.gens[kind]
|
||||
result = gen(img, iter_i=iter_i, raw_image=raw_image)
|
||||
if self.invert_proba > 0 and random.random() < self.invert_proba:
|
||||
result = 1 - result
|
||||
return result
|
||||
|
||||
|
||||
@ANNOTATORS.register_class()
|
||||
class InpaintingAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
para_dict = {}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.mask_cfg = cfg.get(
|
||||
'MASK_CFG', {
|
||||
'irregular_proba': 0.5,
|
||||
'irregular_kwargs': {
|
||||
'min_times': 4,
|
||||
'max_times': 10,
|
||||
'max_width': 100,
|
||||
'max_angle': 4,
|
||||
'max_len': 200
|
||||
},
|
||||
'box_proba': 0.5,
|
||||
'box_kwargs': {
|
||||
'margin': 0,
|
||||
'bbox_min_size': 30,
|
||||
'bbox_max_size': 150,
|
||||
'max_times': 5,
|
||||
'min_times': 1
|
||||
}
|
||||
})
|
||||
self.mask_cfg = Config.get_dict(self.mask_cfg) if not isinstance(
|
||||
self.mask_cfg, dict) else self.mask_cfg
|
||||
self.mask_generator = MixedMaskGenerator(**self.mask_cfg)
|
||||
self.return_mask = cfg.get('RETURN_MASK', False)
|
||||
self.return_invert = cfg.get('RETURN_INVERT', True)
|
||||
self.mask_color = cfg.get('MASK_COLOR', 0)
|
||||
|
||||
def forward(self, image, mask=None, return_mask=None, return_invert=None, mask_color=None):
|
||||
return_mask = return_mask if return_mask is not None else self.return_mask
|
||||
return_invert = return_invert if return_invert is not None else self.return_invert
|
||||
mask_color = mask_color if mask_color is not None else self.mask_color
|
||||
if isinstance(image, Image.Image):
|
||||
image = np.array(image)
|
||||
elif isinstance(image, torch.Tensor):
|
||||
image = image.detach().cpu().numpy()
|
||||
elif isinstance(image, np.ndarray):
|
||||
image = image.copy()
|
||||
else:
|
||||
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
|
||||
|
||||
if mask is not None:
|
||||
if mask_color:
|
||||
image[np.array(mask) == 255] = mask_color
|
||||
else:
|
||||
image[np.array(mask) == 255] = 0
|
||||
else:
|
||||
img = np.transpose(image, (2, 0, 1))
|
||||
mask = self.mask_generator(img)
|
||||
mask = (np.transpose(mask, (1, 2, 0)).squeeze(-1) * 255).astype(np.uint8)
|
||||
if return_invert:
|
||||
mask = invert_image(mask)
|
||||
colored_mask = np.zeros_like(image)
|
||||
if mask_color: colored_mask[:] = mask_color
|
||||
image = np.where(mask[:, :, np.newaxis] == 255, colored_mask, image)
|
||||
|
||||
if return_mask:
|
||||
ret_data = {
|
||||
'image': np.array(image),
|
||||
'mask': np.array(mask)
|
||||
}
|
||||
else:
|
||||
ret_data = np.array(image)
|
||||
return ret_data
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('ANNOTATORS',
|
||||
__class__.__name__,
|
||||
InpaintingAnnotator.para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,98 @@
|
||||
from abc import ABCMeta
|
||||
|
||||
import torch
|
||||
import cv2
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||
from scepter.modules.annotator.registry import ANNOTATORS
|
||||
|
||||
def dilate_mask(mask, dilate_factor=15):
|
||||
mask = mask.astype(np.uint8)
|
||||
mask = cv2.dilate(
|
||||
mask,
|
||||
np.ones((dilate_factor, dilate_factor), np.uint8),
|
||||
iterations=1
|
||||
)
|
||||
return mask
|
||||
|
||||
@ANNOTATORS.register_class()
|
||||
class LamaAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
para_dict = {}
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
from modelscope.pipelines.builder import PIPELINES
|
||||
from modelscope.pipelines.cv import ImageInpaintingPipeline
|
||||
from modelscope.pipelines import pipeline
|
||||
from modelscope.utils.constant import Tasks
|
||||
from modelscope.metainfo import Pipelines
|
||||
from modelscope.models.cv.image_inpainting.refinement import refine_predict
|
||||
from torch.utils.data._utils.collate import default_collate
|
||||
|
||||
@PIPELINES.register_module(Tasks.image_inpainting, module_name=Pipelines.image_inpainting + "-v2")
|
||||
class ImageInpaintingPipelineV2(ImageInpaintingPipeline):
|
||||
def perform_inference(self, data):
|
||||
px_budget = 9000000
|
||||
batch = default_collate([data])
|
||||
if self.refine:
|
||||
assert 'unpad_to_size' in batch, 'Unpadded size is required for the refinement'
|
||||
assert 'cuda' in str(self.device), 'GPU is required for refinement'
|
||||
gpu_ids = str(self.device).split(':')[-1]
|
||||
cur_res = refine_predict(
|
||||
batch,
|
||||
self.infer_model,
|
||||
gpu_ids=gpu_ids,
|
||||
modulo=self.pad_out_to_modulo,
|
||||
n_iters=15,
|
||||
lr=0.002,
|
||||
min_side=512,
|
||||
max_scales=3,
|
||||
px_budget=px_budget)
|
||||
cur_res = cur_res[0].permute(1, 2, 0).detach().cpu().numpy()
|
||||
else:
|
||||
with torch.no_grad():
|
||||
batch = self.move_to_device(batch, self.device)
|
||||
batch['mask'] = (batch['mask'] > 0) * 1
|
||||
batch = self.infer_model(batch)
|
||||
cur_res = batch['inpainted'][0].permute(
|
||||
1, 2, 0).detach().cpu().numpy()
|
||||
unpad_to_size = batch.get('unpad_to_size', None)
|
||||
if unpad_to_size is not None:
|
||||
orig_height, orig_width = unpad_to_size
|
||||
cur_res = cur_res[:orig_height, :orig_width]
|
||||
|
||||
cur_res = np.clip(cur_res * 255, 0, 255).astype('uint8')
|
||||
cur_res = cv2.cvtColor(cur_res, cv2.COLOR_RGB2BGR)
|
||||
return cur_res
|
||||
|
||||
lama_model_dir = FS.get_dir_to_local_dir(cfg.PRETRAINED_MODEL)
|
||||
self.lama_model = pipeline(Tasks.image_inpainting, model=lama_model_dir,
|
||||
pipeline_name=Pipelines.image_inpainting + "-v2", refine=True,
|
||||
device="cuda:{}".format(we.device_id))
|
||||
def forward(self, image, mask):
|
||||
mask = dilate_mask(mask, dilate_factor=19)
|
||||
input_mask = Image.fromarray(mask)
|
||||
mask_expanded = np.tile(np.expand_dims(mask, axis=-1), (1, 1, 3))
|
||||
input_image_np = np.array(image)
|
||||
input_image_np[mask_expanded == 255] = 0
|
||||
input_image = Image.fromarray(input_image_np)
|
||||
input = {
|
||||
'img': input_image,
|
||||
'mask': input_mask,
|
||||
}
|
||||
result = self.lama_model(input)
|
||||
output_img = result['output_img']
|
||||
return output_img[..., ::-1]
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('ANNOTATORS',
|
||||
__class__.__name__,
|
||||
LamaAnnotator.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
|
||||
@@ -9,9 +9,10 @@ import os
|
||||
from abc import ABCMeta
|
||||
from collections import OrderedDict
|
||||
|
||||
import numpy as np
|
||||
|
||||
import cv2
|
||||
import matplotlib
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from PIL import Image
|
||||
@@ -478,7 +479,7 @@ class Hand(object):
|
||||
map_ori = heatmap_avg[:, :, part]
|
||||
one_heatmap = gaussian_filter(map_ori, sigma=3)
|
||||
binary = np.ascontiguousarray(one_heatmap > thre, dtype=np.uint8)
|
||||
# 全部小于阈值
|
||||
|
||||
if np.sum(binary) == 0:
|
||||
all_peaks.append([0, 0])
|
||||
continue
|
||||
|
||||
@@ -0,0 +1,189 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import math
|
||||
import random
|
||||
from abc import ABCMeta
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image, ImageDraw
|
||||
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||
from scepter.modules.annotator.registry import ANNOTATORS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
|
||||
|
||||
@ANNOTATORS.register_class()
|
||||
class OutpaintingAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
para_dict = {}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.mask_blur = cfg.get('MASK_BLUR', 0)
|
||||
self.random_cfg = cfg.get('RANDOM_CFG', None)
|
||||
self.return_mask = cfg.get('RETURN_MASK', False)
|
||||
self.return_source = cfg.get('RETURN_SOURCE', True)
|
||||
self.keep_padding_ratio = cfg.get('KEEP_PADDING_RATIO', 64)
|
||||
self.mask_color = cfg.get('MASK_COLOR', 0)
|
||||
|
||||
def get_box(self, mask):
|
||||
locs = np.where(mask == 255)
|
||||
if len(locs) < 1 or locs[0].shape[0] < 1 or locs[1].shape[0] < 1:
|
||||
return None
|
||||
left, right = np.min(locs[1]), np.max(locs[1])
|
||||
top, bottom = np.min(locs[0]), np.max(locs[0])
|
||||
return [left, top, right, bottom]
|
||||
|
||||
def forward(self,
|
||||
image,
|
||||
ratio=0.3,
|
||||
mask=None,
|
||||
direction=['left', 'right', 'up', 'down'],
|
||||
return_mask=None,
|
||||
return_source=None,
|
||||
mask_color=None):
|
||||
return_mask = return_mask if return_mask is not None else self.return_mask
|
||||
return_source = return_source if return_source is not None else self.return_source
|
||||
mask_color = mask_color if mask_color is not None else self.mask_color
|
||||
if isinstance(image, Image.Image):
|
||||
image = image
|
||||
elif isinstance(image, torch.Tensor):
|
||||
image = Image.fromarray(image.detach().cpu().numpy())
|
||||
elif isinstance(image, np.ndarray):
|
||||
image = Image.fromarray(image.copy())
|
||||
else:
|
||||
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
|
||||
|
||||
if self.random_cfg:
|
||||
direction_range = self.random_cfg.get(
|
||||
'DIRECTION_RANGE', ['left', 'right', 'up', 'down'])
|
||||
ratio_range = self.random_cfg.get('RATIO_RANGE', [0.0, 1.0])
|
||||
direction = random.sample(
|
||||
direction_range,
|
||||
random.choice(list(range(1,
|
||||
len(direction_range) + 1))))
|
||||
ratio = random.uniform(ratio_range[0], ratio_range[1])
|
||||
|
||||
if mask is None:
|
||||
init_image = image
|
||||
src_width, src_height = init_image.width, init_image.height
|
||||
left = int(ratio * src_width) if 'left' in direction else 0
|
||||
right = int(ratio * src_width) if 'right' in direction else 0
|
||||
up = int(ratio * src_height) if 'up' in direction else 0
|
||||
down = int(ratio * src_height) if 'down' in direction else 0
|
||||
# print(direction, ratio, left, right, up, down)
|
||||
tar_width = math.ceil(
|
||||
(src_width + left + right) /
|
||||
self.keep_padding_ratio) * self.keep_padding_ratio
|
||||
tar_height = math.ceil(
|
||||
(src_height + up + down) /
|
||||
self.keep_padding_ratio) * self.keep_padding_ratio
|
||||
if left > 0:
|
||||
left = left * (tar_width - src_width) // (left + right)
|
||||
if right > 0:
|
||||
right = tar_width - src_width - left
|
||||
if up > 0:
|
||||
up = up * (tar_height - src_height) // (up + down)
|
||||
if down > 0:
|
||||
down = tar_height - src_height - up
|
||||
if mask_color is not None:
|
||||
img = Image.new('RGB', (tar_width, tar_height), color=mask_color)
|
||||
else:
|
||||
img = Image.new('RGB', (tar_width, tar_height))
|
||||
img.paste(init_image, (left, up))
|
||||
mask = Image.new('L', (img.width, img.height), 'white')
|
||||
draw = ImageDraw.Draw(mask)
|
||||
|
||||
draw.rectangle(
|
||||
(left + (self.mask_blur * 2 if left > 0 else 0), up +
|
||||
(self.mask_blur * 2 if up > 0 else 0), mask.width - right -
|
||||
(self.mask_blur * 2 if right > 0 else 0), mask.height - down -
|
||||
(self.mask_blur * 2 if down > 0 else 0)),
|
||||
fill='black')
|
||||
else:
|
||||
bbox = self.get_box(np.array(mask))
|
||||
if bbox is None:
|
||||
img = image
|
||||
mask = mask
|
||||
init_image = image
|
||||
else:
|
||||
mask = Image.new('L', (image.width, image.height), 'white')
|
||||
mask_zero = Image.new('L', (bbox[2]-bbox[0], bbox[3]-bbox[1]), 'black')
|
||||
mask.paste(mask_zero, (bbox[0], bbox[1]))
|
||||
crop_image = image.crop(bbox)
|
||||
init_image = Image.new('RGB', (image.width, image.height), 'black')
|
||||
init_image.paste(crop_image, (bbox[0], bbox[1]))
|
||||
img = image
|
||||
if return_mask:
|
||||
if return_source:
|
||||
ret_data = {'src_image': np.array(init_image), 'image': np.array(img), 'mask': np.array(mask)}
|
||||
else:
|
||||
ret_data = {'image': np.array(img), 'mask': np.array(mask)}
|
||||
else:
|
||||
if return_source:
|
||||
ret_data = {'src_image': np.array(init_image), 'image': np.array(img)}
|
||||
else:
|
||||
ret_data = np.array(img)
|
||||
return ret_data
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('ANNOTATORS',
|
||||
__class__.__name__,
|
||||
OutpaintingAnnotator.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@ANNOTATORS.register_class()
|
||||
class OutpaintingResize(BaseAnnotator, metaclass=ABCMeta):
|
||||
para_dict = {}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
|
||||
def get_box(self, mask):
|
||||
locs = np.where(mask == 0)
|
||||
if len(locs) < 1 or locs[0].shape[0] < 1 or locs[1].shape[0] < 1:
|
||||
return None
|
||||
left, right = np.min(locs[1]), np.max(locs[1])
|
||||
top, bottom = np.min(locs[0]), np.max(locs[0])
|
||||
return [left, top, right, bottom]
|
||||
|
||||
def forward(self,
|
||||
image,
|
||||
target_image,
|
||||
mask=None
|
||||
):
|
||||
if isinstance(image, Image.Image):
|
||||
image = image
|
||||
elif isinstance(image, torch.Tensor):
|
||||
image = Image.fromarray(image.detach().cpu().numpy())
|
||||
elif isinstance(image, np.ndarray):
|
||||
image = Image.fromarray(image.copy())
|
||||
else:
|
||||
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
|
||||
|
||||
if isinstance(target_image, Image.Image):
|
||||
target_image = target_image
|
||||
elif isinstance(target_image, torch.Tensor):
|
||||
target_image = Image.fromarray(target_image.detach().cpu().numpy())
|
||||
elif isinstance(target_image, np.ndarray):
|
||||
target_image = Image.fromarray(target_image.copy())
|
||||
else:
|
||||
raise f'Unsurpport datatype{type(target_image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
|
||||
|
||||
bbox = self.get_box(np.array(mask))
|
||||
if bbox is None:
|
||||
init_image = image
|
||||
else:
|
||||
paste_img = image.resize((bbox[2]-bbox[0], bbox[3]-bbox[1]))
|
||||
init_image = Image.new('RGB', (target_image.width, target_image.height), 'black')
|
||||
init_image.paste(paste_img, (bbox[0], bbox[1]))
|
||||
ret_data = {'src_image': np.array(init_image)}
|
||||
return ret_data
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('ANNOTATORS',
|
||||
__class__.__name__,
|
||||
OutpaintingResize.para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,935 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
import math
|
||||
from abc import ABCMeta
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torchvision.transforms as TT
|
||||
from einops import rearrange
|
||||
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||
from scepter.modules.annotator.registry import ANNOTATORS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
CONFIGS = {
|
||||
'baseline': {
|
||||
'layer0': 'cv',
|
||||
'layer1': 'cv',
|
||||
'layer2': 'cv',
|
||||
'layer3': 'cv',
|
||||
'layer4': 'cv',
|
||||
'layer5': 'cv',
|
||||
'layer6': 'cv',
|
||||
'layer7': 'cv',
|
||||
'layer8': 'cv',
|
||||
'layer9': 'cv',
|
||||
'layer10': 'cv',
|
||||
'layer11': 'cv',
|
||||
'layer12': 'cv',
|
||||
'layer13': 'cv',
|
||||
'layer14': 'cv',
|
||||
'layer15': 'cv',
|
||||
},
|
||||
'c-v15': {
|
||||
'layer0': 'cd',
|
||||
'layer1': 'cv',
|
||||
'layer2': 'cv',
|
||||
'layer3': 'cv',
|
||||
'layer4': 'cv',
|
||||
'layer5': 'cv',
|
||||
'layer6': 'cv',
|
||||
'layer7': 'cv',
|
||||
'layer8': 'cv',
|
||||
'layer9': 'cv',
|
||||
'layer10': 'cv',
|
||||
'layer11': 'cv',
|
||||
'layer12': 'cv',
|
||||
'layer13': 'cv',
|
||||
'layer14': 'cv',
|
||||
'layer15': 'cv',
|
||||
},
|
||||
'a-v15': {
|
||||
'layer0': 'ad',
|
||||
'layer1': 'cv',
|
||||
'layer2': 'cv',
|
||||
'layer3': 'cv',
|
||||
'layer4': 'cv',
|
||||
'layer5': 'cv',
|
||||
'layer6': 'cv',
|
||||
'layer7': 'cv',
|
||||
'layer8': 'cv',
|
||||
'layer9': 'cv',
|
||||
'layer10': 'cv',
|
||||
'layer11': 'cv',
|
||||
'layer12': 'cv',
|
||||
'layer13': 'cv',
|
||||
'layer14': 'cv',
|
||||
'layer15': 'cv',
|
||||
},
|
||||
'r-v15': {
|
||||
'layer0': 'rd',
|
||||
'layer1': 'cv',
|
||||
'layer2': 'cv',
|
||||
'layer3': 'cv',
|
||||
'layer4': 'cv',
|
||||
'layer5': 'cv',
|
||||
'layer6': 'cv',
|
||||
'layer7': 'cv',
|
||||
'layer8': 'cv',
|
||||
'layer9': 'cv',
|
||||
'layer10': 'cv',
|
||||
'layer11': 'cv',
|
||||
'layer12': 'cv',
|
||||
'layer13': 'cv',
|
||||
'layer14': 'cv',
|
||||
'layer15': 'cv',
|
||||
},
|
||||
'cvvv4': {
|
||||
'layer0': 'cd',
|
||||
'layer1': 'cv',
|
||||
'layer2': 'cv',
|
||||
'layer3': 'cv',
|
||||
'layer4': 'cd',
|
||||
'layer5': 'cv',
|
||||
'layer6': 'cv',
|
||||
'layer7': 'cv',
|
||||
'layer8': 'cd',
|
||||
'layer9': 'cv',
|
||||
'layer10': 'cv',
|
||||
'layer11': 'cv',
|
||||
'layer12': 'cd',
|
||||
'layer13': 'cv',
|
||||
'layer14': 'cv',
|
||||
'layer15': 'cv',
|
||||
},
|
||||
'avvv4': {
|
||||
'layer0': 'ad',
|
||||
'layer1': 'cv',
|
||||
'layer2': 'cv',
|
||||
'layer3': 'cv',
|
||||
'layer4': 'ad',
|
||||
'layer5': 'cv',
|
||||
'layer6': 'cv',
|
||||
'layer7': 'cv',
|
||||
'layer8': 'ad',
|
||||
'layer9': 'cv',
|
||||
'layer10': 'cv',
|
||||
'layer11': 'cv',
|
||||
'layer12': 'ad',
|
||||
'layer13': 'cv',
|
||||
'layer14': 'cv',
|
||||
'layer15': 'cv',
|
||||
},
|
||||
'rvvv4': {
|
||||
'layer0': 'rd',
|
||||
'layer1': 'cv',
|
||||
'layer2': 'cv',
|
||||
'layer3': 'cv',
|
||||
'layer4': 'rd',
|
||||
'layer5': 'cv',
|
||||
'layer6': 'cv',
|
||||
'layer7': 'cv',
|
||||
'layer8': 'rd',
|
||||
'layer9': 'cv',
|
||||
'layer10': 'cv',
|
||||
'layer11': 'cv',
|
||||
'layer12': 'rd',
|
||||
'layer13': 'cv',
|
||||
'layer14': 'cv',
|
||||
'layer15': 'cv',
|
||||
},
|
||||
'cccv4': {
|
||||
'layer0': 'cd',
|
||||
'layer1': 'cd',
|
||||
'layer2': 'cd',
|
||||
'layer3': 'cv',
|
||||
'layer4': 'cd',
|
||||
'layer5': 'cd',
|
||||
'layer6': 'cd',
|
||||
'layer7': 'cv',
|
||||
'layer8': 'cd',
|
||||
'layer9': 'cd',
|
||||
'layer10': 'cd',
|
||||
'layer11': 'cv',
|
||||
'layer12': 'cd',
|
||||
'layer13': 'cd',
|
||||
'layer14': 'cd',
|
||||
'layer15': 'cv',
|
||||
},
|
||||
'aaav4': {
|
||||
'layer0': 'ad',
|
||||
'layer1': 'ad',
|
||||
'layer2': 'ad',
|
||||
'layer3': 'cv',
|
||||
'layer4': 'ad',
|
||||
'layer5': 'ad',
|
||||
'layer6': 'ad',
|
||||
'layer7': 'cv',
|
||||
'layer8': 'ad',
|
||||
'layer9': 'ad',
|
||||
'layer10': 'ad',
|
||||
'layer11': 'cv',
|
||||
'layer12': 'ad',
|
||||
'layer13': 'ad',
|
||||
'layer14': 'ad',
|
||||
'layer15': 'cv',
|
||||
},
|
||||
'rrrv4': {
|
||||
'layer0': 'rd',
|
||||
'layer1': 'rd',
|
||||
'layer2': 'rd',
|
||||
'layer3': 'cv',
|
||||
'layer4': 'rd',
|
||||
'layer5': 'rd',
|
||||
'layer6': 'rd',
|
||||
'layer7': 'cv',
|
||||
'layer8': 'rd',
|
||||
'layer9': 'rd',
|
||||
'layer10': 'rd',
|
||||
'layer11': 'cv',
|
||||
'layer12': 'rd',
|
||||
'layer13': 'rd',
|
||||
'layer14': 'rd',
|
||||
'layer15': 'cv',
|
||||
},
|
||||
'c16': {
|
||||
'layer0': 'cd',
|
||||
'layer1': 'cd',
|
||||
'layer2': 'cd',
|
||||
'layer3': 'cd',
|
||||
'layer4': 'cd',
|
||||
'layer5': 'cd',
|
||||
'layer6': 'cd',
|
||||
'layer7': 'cd',
|
||||
'layer8': 'cd',
|
||||
'layer9': 'cd',
|
||||
'layer10': 'cd',
|
||||
'layer11': 'cd',
|
||||
'layer12': 'cd',
|
||||
'layer13': 'cd',
|
||||
'layer14': 'cd',
|
||||
'layer15': 'cd',
|
||||
},
|
||||
'a16': {
|
||||
'layer0': 'ad',
|
||||
'layer1': 'ad',
|
||||
'layer2': 'ad',
|
||||
'layer3': 'ad',
|
||||
'layer4': 'ad',
|
||||
'layer5': 'ad',
|
||||
'layer6': 'ad',
|
||||
'layer7': 'ad',
|
||||
'layer8': 'ad',
|
||||
'layer9': 'ad',
|
||||
'layer10': 'ad',
|
||||
'layer11': 'ad',
|
||||
'layer12': 'ad',
|
||||
'layer13': 'ad',
|
||||
'layer14': 'ad',
|
||||
'layer15': 'ad',
|
||||
},
|
||||
'r16': {
|
||||
'layer0': 'rd',
|
||||
'layer1': 'rd',
|
||||
'layer2': 'rd',
|
||||
'layer3': 'rd',
|
||||
'layer4': 'rd',
|
||||
'layer5': 'rd',
|
||||
'layer6': 'rd',
|
||||
'layer7': 'rd',
|
||||
'layer8': 'rd',
|
||||
'layer9': 'rd',
|
||||
'layer10': 'rd',
|
||||
'layer11': 'rd',
|
||||
'layer12': 'rd',
|
||||
'layer13': 'rd',
|
||||
'layer14': 'rd',
|
||||
'layer15': 'rd',
|
||||
},
|
||||
'carv4': {
|
||||
'layer0': 'cd',
|
||||
'layer1': 'ad',
|
||||
'layer2': 'rd',
|
||||
'layer3': 'cv',
|
||||
'layer4': 'cd',
|
||||
'layer5': 'ad',
|
||||
'layer6': 'rd',
|
||||
'layer7': 'cv',
|
||||
'layer8': 'cd',
|
||||
'layer9': 'ad',
|
||||
'layer10': 'rd',
|
||||
'layer11': 'cv',
|
||||
'layer12': 'cd',
|
||||
'layer13': 'ad',
|
||||
'layer14': 'rd',
|
||||
'layer15': 'cv'
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
def create_conv_func(op_type):
|
||||
assert op_type in ['cv', 'cd', 'ad',
|
||||
'rd'], 'unknown op type: %s' % str(op_type)
|
||||
if op_type == 'cv':
|
||||
return F.conv2d
|
||||
if op_type == 'cd':
|
||||
|
||||
def func(x,
|
||||
weights,
|
||||
bias=None,
|
||||
stride=1,
|
||||
padding=0,
|
||||
dilation=1,
|
||||
groups=1):
|
||||
assert dilation in [1,
|
||||
2], 'dilation for cd_conv should be in 1 or 2'
|
||||
assert weights.size(2) == 3 and weights.size(3) == 3, \
|
||||
'kernel size for cd_conv should be 3x3'
|
||||
assert padding == dilation, 'padding for cd_conv set wrong'
|
||||
|
||||
weights_c = weights.sum(dim=[2, 3], keepdim=True)
|
||||
yc = F.conv2d(x,
|
||||
weights_c,
|
||||
stride=stride,
|
||||
padding=0,
|
||||
groups=groups)
|
||||
y = F.conv2d(x,
|
||||
weights,
|
||||
bias,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
dilation=dilation,
|
||||
groups=groups)
|
||||
return y - yc
|
||||
|
||||
return func
|
||||
elif op_type == 'ad':
|
||||
|
||||
def func(x,
|
||||
weights,
|
||||
bias=None,
|
||||
stride=1,
|
||||
padding=0,
|
||||
dilation=1,
|
||||
groups=1):
|
||||
assert dilation in [1,
|
||||
2], 'dilation for ad_conv should be in 1 or 2'
|
||||
assert weights.size(2) == 3 and weights.size(3) == 3, \
|
||||
'kernel size for ad_conv should be 3x3'
|
||||
assert padding == dilation, 'padding for ad_conv set wrong'
|
||||
|
||||
shape = weights.shape
|
||||
weights = weights.view(shape[0], shape[1], -1)
|
||||
# clock-wise
|
||||
weights_conv = (
|
||||
weights -
|
||||
weights[:, :, [3, 0, 1, 6, 4, 2, 7, 8, 5]]).view(shape)
|
||||
y = F.conv2d(x,
|
||||
weights_conv,
|
||||
bias,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
dilation=dilation,
|
||||
groups=groups)
|
||||
return y
|
||||
|
||||
return func
|
||||
elif op_type == 'rd':
|
||||
|
||||
def func(x,
|
||||
weights,
|
||||
bias=None,
|
||||
stride=1,
|
||||
padding=0,
|
||||
dilation=1,
|
||||
groups=1):
|
||||
assert dilation in [1,
|
||||
2], 'dilation for rd_conv should be in 1 or 2'
|
||||
assert weights.size(2) == 3 and weights.size(3) == 3, \
|
||||
'kernel size for rd_conv should be 3x3'
|
||||
padding = 2 * dilation
|
||||
|
||||
shape = weights.shape
|
||||
if weights.is_cuda:
|
||||
buffer = torch.cuda.FloatTensor(shape[0], shape[1],
|
||||
5 * 5).fill_(0)
|
||||
else:
|
||||
buffer = torch.zeros(shape[0], shape[1], 5 * 5)
|
||||
weights = weights.view(shape[0], shape[1], -1)
|
||||
buffer[:, :, [0, 2, 4, 10, 14, 20, 22, 24]] = weights[:, :, 1:]
|
||||
buffer[:, :, [6, 7, 8, 11, 13, 16, 17, 18]] = -weights[:, :, 1:]
|
||||
buffer[:, :, 12] = 0
|
||||
buffer = buffer.view(shape[0], shape[1], 5, 5)
|
||||
y = F.conv2d(x,
|
||||
buffer,
|
||||
bias,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
dilation=dilation,
|
||||
groups=groups)
|
||||
return y
|
||||
|
||||
return func
|
||||
else:
|
||||
print('impossible to be here unless you force that', flush=True)
|
||||
return None
|
||||
|
||||
|
||||
def config_model(model):
|
||||
model_options = list(CONFIGS.keys())
|
||||
assert model in model_options, \
|
||||
'unrecognized model, please choose from %s' % str(model_options)
|
||||
|
||||
pdcs = []
|
||||
for i in range(16):
|
||||
layer_name = 'layer%d' % i
|
||||
op = CONFIGS[model][layer_name]
|
||||
pdcs.append(create_conv_func(op))
|
||||
return pdcs
|
||||
|
||||
|
||||
def config_model_converted(model):
|
||||
model_options = list(CONFIGS.keys())
|
||||
assert model in model_options, \
|
||||
'unrecognized model, please choose from %s' % str(model_options)
|
||||
|
||||
pdcs = []
|
||||
for i in range(16):
|
||||
layer_name = 'layer%d' % i
|
||||
op = CONFIGS[model][layer_name]
|
||||
pdcs.append(op)
|
||||
return pdcs
|
||||
|
||||
|
||||
def convert_pdc(op, weight):
|
||||
if op == 'cv':
|
||||
return weight
|
||||
elif op == 'cd':
|
||||
shape = weight.shape
|
||||
weight_c = weight.sum(dim=[2, 3])
|
||||
weight = weight.view(shape[0], shape[1], -1)
|
||||
weight[:, :, 4] = weight[:, :, 4] - weight_c
|
||||
weight = weight.view(shape)
|
||||
return weight
|
||||
elif op == 'ad':
|
||||
shape = weight.shape
|
||||
weight = weight.view(shape[0], shape[1], -1)
|
||||
weight_conv = (weight -
|
||||
weight[:, :, [3, 0, 1, 6, 4, 2, 7, 8, 5]]).view(shape)
|
||||
return weight_conv
|
||||
elif op == 'rd':
|
||||
shape = weight.shape
|
||||
buffer = torch.zeros(shape[0], shape[1], 5 * 5, device=weight.device)
|
||||
weight = weight.view(shape[0], shape[1], -1)
|
||||
buffer[:, :, [0, 2, 4, 10, 14, 20, 22, 24]] = weight[:, :, 1:]
|
||||
buffer[:, :, [6, 7, 8, 11, 13, 16, 17, 18]] = -weight[:, :, 1:]
|
||||
buffer = buffer.view(shape[0], shape[1], 5, 5)
|
||||
return buffer
|
||||
raise ValueError('wrong op {}'.format(str(op)))
|
||||
|
||||
|
||||
def convert_pidinet(state_dict, config):
|
||||
pdcs = config_model_converted(config)
|
||||
new_dict = {}
|
||||
for pname, p in state_dict.items():
|
||||
if 'init_block.weight' in pname:
|
||||
new_dict[pname] = convert_pdc(pdcs[0], p)
|
||||
elif 'block1_1.conv1.weight' in pname:
|
||||
new_dict[pname] = convert_pdc(pdcs[1], p)
|
||||
elif 'block1_2.conv1.weight' in pname:
|
||||
new_dict[pname] = convert_pdc(pdcs[2], p)
|
||||
elif 'block1_3.conv1.weight' in pname:
|
||||
new_dict[pname] = convert_pdc(pdcs[3], p)
|
||||
elif 'block2_1.conv1.weight' in pname:
|
||||
new_dict[pname] = convert_pdc(pdcs[4], p)
|
||||
elif 'block2_2.conv1.weight' in pname:
|
||||
new_dict[pname] = convert_pdc(pdcs[5], p)
|
||||
elif 'block2_3.conv1.weight' in pname:
|
||||
new_dict[pname] = convert_pdc(pdcs[6], p)
|
||||
elif 'block2_4.conv1.weight' in pname:
|
||||
new_dict[pname] = convert_pdc(pdcs[7], p)
|
||||
elif 'block3_1.conv1.weight' in pname:
|
||||
new_dict[pname] = convert_pdc(pdcs[8], p)
|
||||
elif 'block3_2.conv1.weight' in pname:
|
||||
new_dict[pname] = convert_pdc(pdcs[9], p)
|
||||
elif 'block3_3.conv1.weight' in pname:
|
||||
new_dict[pname] = convert_pdc(pdcs[10], p)
|
||||
elif 'block3_4.conv1.weight' in pname:
|
||||
new_dict[pname] = convert_pdc(pdcs[11], p)
|
||||
elif 'block4_1.conv1.weight' in pname:
|
||||
new_dict[pname] = convert_pdc(pdcs[12], p)
|
||||
elif 'block4_2.conv1.weight' in pname:
|
||||
new_dict[pname] = convert_pdc(pdcs[13], p)
|
||||
elif 'block4_3.conv1.weight' in pname:
|
||||
new_dict[pname] = convert_pdc(pdcs[14], p)
|
||||
elif 'block4_4.conv1.weight' in pname:
|
||||
new_dict[pname] = convert_pdc(pdcs[15], p)
|
||||
else:
|
||||
new_dict[pname] = p
|
||||
return new_dict
|
||||
|
||||
|
||||
class Conv2d(nn.Module):
|
||||
def __init__(self,
|
||||
pdc,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride=1,
|
||||
padding=0,
|
||||
dilation=1,
|
||||
groups=1,
|
||||
bias=False):
|
||||
super().__init__()
|
||||
if in_channels % groups != 0:
|
||||
raise ValueError('in_channels must be divisible by groups')
|
||||
if out_channels % groups != 0:
|
||||
raise ValueError('out_channels must be divisible by groups')
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = out_channels
|
||||
self.kernel_size = kernel_size
|
||||
self.stride = stride
|
||||
self.padding = padding
|
||||
self.dilation = dilation
|
||||
self.groups = groups
|
||||
self.weight = nn.Parameter(
|
||||
torch.Tensor(out_channels, in_channels // groups, kernel_size,
|
||||
kernel_size))
|
||||
if bias:
|
||||
self.bias = nn.Parameter(torch.Tensor(out_channels))
|
||||
else:
|
||||
self.register_parameter('bias', None)
|
||||
self.reset_parameters()
|
||||
self.pdc = pdc
|
||||
|
||||
def reset_parameters(self):
|
||||
nn.init.kaiming_uniform_(self.weight, a=math.sqrt(5))
|
||||
if self.bias is not None:
|
||||
fan_in, _ = nn.init._calculate_fan_in_and_fan_out(self.weight)
|
||||
bound = 1 / math.sqrt(fan_in)
|
||||
nn.init.uniform_(self.bias, -bound, bound)
|
||||
|
||||
def forward(self, input):
|
||||
return self.pdc(input, self.weight, self.bias, self.stride,
|
||||
self.padding, self.dilation, self.groups)
|
||||
|
||||
|
||||
class CSAM(nn.Module):
|
||||
"""
|
||||
Compact Spatial Attention Module
|
||||
"""
|
||||
def __init__(self, channels):
|
||||
super().__init__()
|
||||
|
||||
mid_channels = 4
|
||||
self.relu1 = nn.ReLU()
|
||||
self.conv1 = nn.Conv2d(channels,
|
||||
mid_channels,
|
||||
kernel_size=1,
|
||||
padding=0)
|
||||
self.conv2 = nn.Conv2d(mid_channels,
|
||||
1,
|
||||
kernel_size=3,
|
||||
padding=1,
|
||||
bias=False)
|
||||
self.sigmoid = nn.Sigmoid()
|
||||
nn.init.constant_(self.conv1.bias, 0)
|
||||
|
||||
def forward(self, x):
|
||||
y = self.relu1(x)
|
||||
y = self.conv1(y)
|
||||
y = self.conv2(y)
|
||||
y = self.sigmoid(y)
|
||||
|
||||
return x * y
|
||||
|
||||
|
||||
class CDCM(nn.Module):
|
||||
"""
|
||||
Compact Dilation Convolution based Module
|
||||
"""
|
||||
def __init__(self, in_channels, out_channels):
|
||||
super().__init__()
|
||||
|
||||
self.relu1 = nn.ReLU()
|
||||
self.conv1 = nn.Conv2d(in_channels,
|
||||
out_channels,
|
||||
kernel_size=1,
|
||||
padding=0)
|
||||
self.conv2_1 = nn.Conv2d(out_channels,
|
||||
out_channels,
|
||||
kernel_size=3,
|
||||
dilation=5,
|
||||
padding=5,
|
||||
bias=False)
|
||||
self.conv2_2 = nn.Conv2d(out_channels,
|
||||
out_channels,
|
||||
kernel_size=3,
|
||||
dilation=7,
|
||||
padding=7,
|
||||
bias=False)
|
||||
self.conv2_3 = nn.Conv2d(out_channels,
|
||||
out_channels,
|
||||
kernel_size=3,
|
||||
dilation=9,
|
||||
padding=9,
|
||||
bias=False)
|
||||
self.conv2_4 = nn.Conv2d(out_channels,
|
||||
out_channels,
|
||||
kernel_size=3,
|
||||
dilation=11,
|
||||
padding=11,
|
||||
bias=False)
|
||||
nn.init.constant_(self.conv1.bias, 0)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.relu1(x)
|
||||
x = self.conv1(x)
|
||||
x1 = self.conv2_1(x)
|
||||
x2 = self.conv2_2(x)
|
||||
x3 = self.conv2_3(x)
|
||||
x4 = self.conv2_4(x)
|
||||
return x1 + x2 + x3 + x4
|
||||
|
||||
|
||||
class MapReduce(nn.Module):
|
||||
"""
|
||||
Reduce feature maps into a single edge map
|
||||
"""
|
||||
def __init__(self, channels):
|
||||
super().__init__()
|
||||
self.conv = nn.Conv2d(channels, 1, kernel_size=1, padding=0)
|
||||
nn.init.constant_(self.conv.bias, 0)
|
||||
|
||||
def forward(self, x):
|
||||
return self.conv(x)
|
||||
|
||||
|
||||
class PDCBlock(nn.Module):
|
||||
def __init__(self, pdc, inplane, ouplane, stride=1):
|
||||
super().__init__()
|
||||
self.stride = stride
|
||||
|
||||
self.stride = stride
|
||||
if self.stride > 1:
|
||||
self.pool = nn.MaxPool2d(kernel_size=2, stride=2)
|
||||
self.shortcut = nn.Conv2d(inplane,
|
||||
ouplane,
|
||||
kernel_size=1,
|
||||
padding=0)
|
||||
self.conv1 = Conv2d(pdc,
|
||||
inplane,
|
||||
inplane,
|
||||
kernel_size=3,
|
||||
padding=1,
|
||||
groups=inplane,
|
||||
bias=False)
|
||||
self.relu2 = nn.ReLU()
|
||||
self.conv2 = nn.Conv2d(inplane,
|
||||
ouplane,
|
||||
kernel_size=1,
|
||||
padding=0,
|
||||
bias=False)
|
||||
|
||||
def forward(self, x):
|
||||
if self.stride > 1:
|
||||
x = self.pool(x)
|
||||
y = self.conv1(x)
|
||||
y = self.relu2(y)
|
||||
y = self.conv2(y)
|
||||
if self.stride > 1:
|
||||
x = self.shortcut(x)
|
||||
y = y + x
|
||||
return y
|
||||
|
||||
|
||||
class PDCBlock_converted(nn.Module):
|
||||
"""
|
||||
CPDC, APDC can be converted to vanilla 3x3 convolution
|
||||
RPDC can be converted to vanilla 5x5 convolution
|
||||
"""
|
||||
def __init__(self, pdc, inplane, ouplane, stride=1):
|
||||
super().__init__()
|
||||
self.stride = stride
|
||||
|
||||
if self.stride > 1:
|
||||
self.pool = nn.MaxPool2d(kernel_size=2, stride=2)
|
||||
self.shortcut = nn.Conv2d(inplane,
|
||||
ouplane,
|
||||
kernel_size=1,
|
||||
padding=0)
|
||||
if pdc == 'rd':
|
||||
self.conv1 = nn.Conv2d(inplane,
|
||||
inplane,
|
||||
kernel_size=5,
|
||||
padding=2,
|
||||
groups=inplane,
|
||||
bias=False)
|
||||
else:
|
||||
self.conv1 = nn.Conv2d(inplane,
|
||||
inplane,
|
||||
kernel_size=3,
|
||||
padding=1,
|
||||
groups=inplane,
|
||||
bias=False)
|
||||
self.relu2 = nn.ReLU()
|
||||
self.conv2 = nn.Conv2d(inplane,
|
||||
ouplane,
|
||||
kernel_size=1,
|
||||
padding=0,
|
||||
bias=False)
|
||||
|
||||
def forward(self, x):
|
||||
if self.stride > 1:
|
||||
x = self.pool(x)
|
||||
y = self.conv1(x)
|
||||
y = self.relu2(y)
|
||||
y = self.conv2(y)
|
||||
if self.stride > 1:
|
||||
x = self.shortcut(x)
|
||||
y = y + x
|
||||
return y
|
||||
|
||||
|
||||
class PiDiNet(nn.Module):
|
||||
def __init__(self,
|
||||
inplane,
|
||||
pdcs,
|
||||
dil=None,
|
||||
sa=False,
|
||||
convert=False,
|
||||
mean=[0.485, 0.456, 0.406],
|
||||
std=[0.229, 0.224, 0.225]):
|
||||
super().__init__()
|
||||
self.sa = sa
|
||||
if dil is not None:
|
||||
assert isinstance(dil, int), 'dil should be an int'
|
||||
self.dil = dil
|
||||
self.mean = mean
|
||||
self.std = std
|
||||
|
||||
self.fuseplanes = []
|
||||
|
||||
self.inplane = inplane
|
||||
if convert:
|
||||
if pdcs[0] == 'rd':
|
||||
init_kernel_size = 5
|
||||
init_padding = 2
|
||||
else:
|
||||
init_kernel_size = 3
|
||||
init_padding = 1
|
||||
self.init_block = nn.Conv2d(3,
|
||||
self.inplane,
|
||||
kernel_size=init_kernel_size,
|
||||
padding=init_padding,
|
||||
bias=False)
|
||||
block_class = PDCBlock_converted
|
||||
else:
|
||||
self.init_block = Conv2d(pdcs[0],
|
||||
3,
|
||||
self.inplane,
|
||||
kernel_size=3,
|
||||
padding=1)
|
||||
block_class = PDCBlock
|
||||
|
||||
self.block1_1 = block_class(pdcs[1], self.inplane, self.inplane)
|
||||
self.block1_2 = block_class(pdcs[2], self.inplane, self.inplane)
|
||||
self.block1_3 = block_class(pdcs[3], self.inplane, self.inplane)
|
||||
self.fuseplanes.append(self.inplane) # C
|
||||
|
||||
inplane = self.inplane
|
||||
self.inplane = self.inplane * 2
|
||||
self.block2_1 = block_class(pdcs[4], inplane, self.inplane, stride=2)
|
||||
self.block2_2 = block_class(pdcs[5], self.inplane, self.inplane)
|
||||
self.block2_3 = block_class(pdcs[6], self.inplane, self.inplane)
|
||||
self.block2_4 = block_class(pdcs[7], self.inplane, self.inplane)
|
||||
self.fuseplanes.append(self.inplane) # 2C
|
||||
|
||||
inplane = self.inplane
|
||||
self.inplane = self.inplane * 2
|
||||
self.block3_1 = block_class(pdcs[8], inplane, self.inplane, stride=2)
|
||||
self.block3_2 = block_class(pdcs[9], self.inplane, self.inplane)
|
||||
self.block3_3 = block_class(pdcs[10], self.inplane, self.inplane)
|
||||
self.block3_4 = block_class(pdcs[11], self.inplane, self.inplane)
|
||||
self.fuseplanes.append(self.inplane) # 4C
|
||||
|
||||
self.block4_1 = block_class(pdcs[12],
|
||||
self.inplane,
|
||||
self.inplane,
|
||||
stride=2)
|
||||
self.block4_2 = block_class(pdcs[13], self.inplane, self.inplane)
|
||||
self.block4_3 = block_class(pdcs[14], self.inplane, self.inplane)
|
||||
self.block4_4 = block_class(pdcs[15], self.inplane, self.inplane)
|
||||
self.fuseplanes.append(self.inplane) # 4C
|
||||
|
||||
self.conv_reduces = nn.ModuleList()
|
||||
if self.sa and self.dil is not None:
|
||||
self.attentions = nn.ModuleList()
|
||||
self.dilations = nn.ModuleList()
|
||||
for i in range(4):
|
||||
self.dilations.append(CDCM(self.fuseplanes[i], self.dil))
|
||||
self.attentions.append(CSAM(self.dil))
|
||||
self.conv_reduces.append(MapReduce(self.dil))
|
||||
elif self.sa:
|
||||
self.attentions = nn.ModuleList()
|
||||
for i in range(4):
|
||||
self.attentions.append(CSAM(self.fuseplanes[i]))
|
||||
self.conv_reduces.append(MapReduce(self.fuseplanes[i]))
|
||||
elif self.dil is not None:
|
||||
self.dilations = nn.ModuleList()
|
||||
for i in range(4):
|
||||
self.dilations.append(CDCM(self.fuseplanes[i], self.dil))
|
||||
self.conv_reduces.append(MapReduce(self.dil))
|
||||
else:
|
||||
for i in range(4):
|
||||
self.conv_reduces.append(MapReduce(self.fuseplanes[i]))
|
||||
|
||||
self.classifier = nn.Conv2d(4, 1, kernel_size=1) # has bias
|
||||
nn.init.constant_(self.classifier.weight, 0.25)
|
||||
nn.init.constant_(self.classifier.bias, 0)
|
||||
|
||||
def get_weights(self):
|
||||
conv_weights = []
|
||||
bn_weights = []
|
||||
relu_weights = []
|
||||
for pname, p in self.named_parameters():
|
||||
if 'bn' in pname:
|
||||
bn_weights.append(p)
|
||||
elif 'relu' in pname:
|
||||
relu_weights.append(p)
|
||||
else:
|
||||
conv_weights.append(p)
|
||||
|
||||
return conv_weights, bn_weights, relu_weights
|
||||
|
||||
def forward(self, x):
|
||||
"""x: [B, 3, H, W] within range [0, 1].
|
||||
"""
|
||||
x = (x - x.new_tensor(self.mean).view(1, -1, 1, 1)) / \
|
||||
x.new_tensor(self.std).view(1, -1, 1, 1)
|
||||
h, w = x.size()[2:]
|
||||
|
||||
x = self.init_block(x)
|
||||
|
||||
x1 = self.block1_1(x)
|
||||
x1 = self.block1_2(x1)
|
||||
x1 = self.block1_3(x1)
|
||||
|
||||
x2 = self.block2_1(x1)
|
||||
x2 = self.block2_2(x2)
|
||||
x2 = self.block2_3(x2)
|
||||
x2 = self.block2_4(x2)
|
||||
|
||||
x3 = self.block3_1(x2)
|
||||
x3 = self.block3_2(x3)
|
||||
x3 = self.block3_3(x3)
|
||||
x3 = self.block3_4(x3)
|
||||
|
||||
x4 = self.block4_1(x3)
|
||||
x4 = self.block4_2(x4)
|
||||
x4 = self.block4_3(x4)
|
||||
x4 = self.block4_4(x4)
|
||||
|
||||
x_fuses = []
|
||||
if self.sa and self.dil is not None:
|
||||
for i, xi in enumerate([x1, x2, x3, x4]):
|
||||
x_fuses.append(self.attentions[i](self.dilations[i](xi)))
|
||||
elif self.sa:
|
||||
for i, xi in enumerate([x1, x2, x3, x4]):
|
||||
x_fuses.append(self.attentions[i](xi))
|
||||
elif self.dil is not None:
|
||||
for i, xi in enumerate([x1, x2, x3, x4]):
|
||||
x_fuses.append(self.dilations[i](xi))
|
||||
else:
|
||||
x_fuses = [x1, x2, x3, x4]
|
||||
|
||||
e1 = self.conv_reduces[0](x_fuses[0])
|
||||
e1 = F.interpolate(e1, (h, w), mode='bilinear', align_corners=False)
|
||||
|
||||
e2 = self.conv_reduces[1](x_fuses[1])
|
||||
e2 = F.interpolate(e2, (h, w), mode='bilinear', align_corners=False)
|
||||
|
||||
e3 = self.conv_reduces[2](x_fuses[2])
|
||||
e3 = F.interpolate(e3, (h, w), mode='bilinear', align_corners=False)
|
||||
|
||||
e4 = self.conv_reduces[3](x_fuses[3])
|
||||
e4 = F.interpolate(e4, (h, w), mode='bilinear', align_corners=False)
|
||||
|
||||
outputs = [e1, e2, e3, e4]
|
||||
output = self.classifier(torch.cat(outputs, dim=1))
|
||||
|
||||
outputs.append(output)
|
||||
outputs = [torch.sigmoid(r) for r in outputs]
|
||||
return outputs[-1]
|
||||
|
||||
|
||||
@ANNOTATORS.register_class()
|
||||
class PiDiAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
para_dict = {}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
pretrained_model = cfg.get('PRETRAINED_MODEL', None)
|
||||
vanilla_cnn = cfg.get('VANILLA_CNN', True)
|
||||
pdcs = config_model_converted(
|
||||
'carv4') if vanilla_cnn else config_model('carv4')
|
||||
self.model = PiDiNet(60, pdcs, dil=24, sa=True,
|
||||
convert=vanilla_cnn).eval()
|
||||
if pretrained_model:
|
||||
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
|
||||
state = torch.load(local_path,
|
||||
map_location='cpu')['state_dict']
|
||||
if vanilla_cnn:
|
||||
state = convert_pidinet(state, 'carv4')
|
||||
state = {
|
||||
k[len('module.'):] if k.startswith('module.') else k: v
|
||||
for k, v in state.items()
|
||||
}
|
||||
self.model.load_state_dict(state)
|
||||
|
||||
@torch.no_grad()
|
||||
@torch.inference_mode()
|
||||
@torch.autocast('cuda', enabled=False)
|
||||
def forward(self, image, return_grayscale=False):
|
||||
is_batch = False if len(image.shape) == 3 else True
|
||||
if isinstance(image, torch.Tensor):
|
||||
if len(image.shape) == 3:
|
||||
image = rearrange(image, 'h w c -> 1 c h w')
|
||||
elif len(image.shape) == 4:
|
||||
image = rearrange(image, 'b h w c -> b c h w')
|
||||
else:
|
||||
raise "Unsurpport input image's shape"
|
||||
elif isinstance(image, np.ndarray):
|
||||
image = torch.from_numpy(image.copy()).float()
|
||||
if len(image.shape) == 3:
|
||||
image = rearrange(image, 'h w c -> 1 c h w')
|
||||
elif len(image.shape) == 4:
|
||||
image = rearrange(image, 'b h w c -> b c h w')
|
||||
else:
|
||||
raise "Unsurpport input image's shape"
|
||||
else:
|
||||
raise "Unsurpport input image's type"
|
||||
image = image.float().div(255)
|
||||
image = image.to(we.device_id)
|
||||
edge = self.model(image)
|
||||
edge = edge.squeeze(dim=1)
|
||||
edge = 255 - (edge * 255.0).clip(0, 255) # return white background
|
||||
edge = edge.cpu().numpy()
|
||||
edge = edge.astype(np.uint8)
|
||||
if not is_batch:
|
||||
edge = edge.squeeze()
|
||||
if not return_grayscale:
|
||||
edge = edge[..., None].repeat(3, -1)
|
||||
return edge
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('ANNOTATORS',
|
||||
__class__.__name__,
|
||||
PiDiAnnotator.para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,382 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import math
|
||||
import random
|
||||
from abc import ABCMeta
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision.transforms as T
|
||||
from PIL import Image
|
||||
from scipy import ndimage
|
||||
from pycocotools import mask as mask_utils
|
||||
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||
from scepter.modules.annotator.registry import ANNOTATORS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from sklearn.cluster import KMeans
|
||||
from torchvision.ops.boxes import batched_nms
|
||||
|
||||
|
||||
def find_dominant_color(image, k=1):
|
||||
pixels = image.reshape((-1, 3))
|
||||
mask = (pixels != [0, 0, 0]).all(axis=1)
|
||||
pixels = pixels[mask]
|
||||
try:
|
||||
kmeans = KMeans(n_clusters=k, n_init='auto')
|
||||
kmeans.fit(pixels)
|
||||
dominant_color = kmeans.cluster_centers_.astype(int)[0]
|
||||
except:
|
||||
dominant_color = np.array([255, 255, 255])
|
||||
return dominant_color
|
||||
|
||||
|
||||
def cv2_resize_crop(image, resize_size, crop_size):
|
||||
resize_height, resize_width = resize_size
|
||||
crop_height, crop_width = crop_size
|
||||
|
||||
resized_image = cv2.resize(image, (resize_width, resize_height))
|
||||
|
||||
center_x, center_y = resize_width // 2, resize_height // 2
|
||||
crop_start_x = max(center_x - crop_width // 2, 0)
|
||||
crop_start_y = max(center_y - crop_height // 2, 0)
|
||||
crop_end_x = crop_start_x + crop_width
|
||||
crop_end_y = crop_start_y + crop_height
|
||||
|
||||
crop_end_x = min(crop_end_x, resize_width)
|
||||
crop_end_y = min(crop_end_y, resize_height)
|
||||
|
||||
center_cropped_image = resized_image[crop_start_y:crop_end_y,
|
||||
crop_start_x:crop_end_x]
|
||||
|
||||
return center_cropped_image
|
||||
|
||||
|
||||
@ANNOTATORS.register_class()
|
||||
class ESAMAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
para_dict = {}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
try:
|
||||
from efficient_sam.efficient_sam import build_efficient_sam
|
||||
from segment_anything.utils.amg import (
|
||||
batched_mask_to_box,
|
||||
calculate_stability_score,
|
||||
mask_to_rle_pytorch,
|
||||
remove_small_regions,
|
||||
rle_to_mask,
|
||||
)
|
||||
except:
|
||||
raise NotImplementedError(
|
||||
f'Please install efficient_sam and segment_anything modules.')
|
||||
|
||||
pretrained_model = cfg.get('PRETRAINED_MODEL', None)
|
||||
if pretrained_model:
|
||||
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
|
||||
self.efficient_sam_module = build_efficient_sam(
|
||||
encoder_patch_embed_dim=384,
|
||||
encoder_num_heads=6,
|
||||
checkpoint=local_path).eval().to(we.device_id)
|
||||
self.GRID_SIZE = cfg.get('GRID_SIZE', 16)
|
||||
self.save_mode = cfg.get('SAVE_MODE', 'P')
|
||||
self.use_dominant_color = cfg.get('USE_DOMINANT_COLOR', False)
|
||||
self.return_mask = cfg.get('RETURN_MASK', False)
|
||||
|
||||
@torch.no_grad()
|
||||
def get_predictions_given_embeddings_and_queries(self, img, points,
|
||||
point_labels, model):
|
||||
from segment_anything.utils.amg import calculate_stability_score
|
||||
predicted_masks, predicted_iou = [], []
|
||||
bs = 128
|
||||
num = int(float(self.GRID_SIZE * self.GRID_SIZE) /
|
||||
bs) if self.GRID_SIZE * self.GRID_SIZE % bs == 0 else int(
|
||||
float(self.GRID_SIZE * self.GRID_SIZE) / bs) + 1
|
||||
for i in range(num):
|
||||
predicted_mask_item, predicted_iou_item = model(
|
||||
img[None, ...], points[:, i * bs:(i + 1) * bs, ...],
|
||||
point_labels[:, i * bs:(i + 1) * bs, :])
|
||||
predicted_masks.append(predicted_mask_item)
|
||||
predicted_iou.append(predicted_iou_item)
|
||||
torch.cuda.empty_cache()
|
||||
# # predicted: torch.Size([1, 1024, 3]) torch.Size([1, 1024, 3, 512, 512])
|
||||
# print('predicted: ', predicted_iou.size(), predicted_masks.size())
|
||||
predicted_masks = torch.cat(predicted_masks, dim=1)
|
||||
predicted_iou = torch.cat(predicted_iou, dim=1)
|
||||
sorted_ids = torch.argsort(predicted_iou, dim=-1, descending=True)
|
||||
predicted_iou_scores = torch.take_along_dim(predicted_iou,
|
||||
sorted_ids,
|
||||
dim=2)
|
||||
predicted_masks = torch.take_along_dim(predicted_masks,
|
||||
sorted_ids[..., None, None],
|
||||
dim=2)
|
||||
predicted_masks = predicted_masks[0]
|
||||
iou = predicted_iou_scores[0, :, 0]
|
||||
index_iou = iou > 0.7
|
||||
iou_ = iou[index_iou]
|
||||
masks = predicted_masks[index_iou]
|
||||
score = calculate_stability_score(masks, 0.0, 1.0)
|
||||
score = score[:, 0]
|
||||
index = score > 0.9
|
||||
masks = masks[index]
|
||||
iou_ = iou_[index]
|
||||
masks = torch.ge(masks, 0.0)
|
||||
return masks, iou_
|
||||
|
||||
def singel_mask_to_rle(self, mask):
|
||||
rle = mask_utils.encode(
|
||||
np.array(mask[:, :, None], order='F', dtype='uint8'))[0]
|
||||
rle['counts'] = rle['counts'].decode('utf-8')
|
||||
return rle
|
||||
|
||||
def process_small_region(self, rles):
|
||||
from segment_anything.utils.amg import rle_to_mask, remove_small_regions, \
|
||||
batched_mask_to_box, mask_to_rle_pytorch
|
||||
new_masks = []
|
||||
scores = []
|
||||
min_area = 100
|
||||
nms_thresh = 0.7
|
||||
for rle in rles:
|
||||
mask = rle_to_mask(rle[0])
|
||||
|
||||
mask, changed = remove_small_regions(mask, min_area, mode='holes')
|
||||
unchanged = not changed
|
||||
mask, changed = remove_small_regions(mask,
|
||||
min_area,
|
||||
mode='islands')
|
||||
unchanged = unchanged and not changed
|
||||
|
||||
new_masks.append(torch.as_tensor(mask).unsqueeze(0))
|
||||
# Give score=0 to changed masks and score=1 to unchanged masks
|
||||
# so NMS will prefer ones that didn't need postprocessing
|
||||
scores.append(float(unchanged))
|
||||
|
||||
# Recalculate boxes and remove any new duplicates
|
||||
masks = torch.cat(new_masks, dim=0).to(we.device_id)
|
||||
boxes = batched_mask_to_box(masks)
|
||||
keep_by_nms = batched_nms(
|
||||
boxes.float(),
|
||||
torch.as_tensor(scores).to(we.device_id),
|
||||
torch.zeros_like(boxes[:, 0]), # categories
|
||||
iou_threshold=nms_thresh,
|
||||
)
|
||||
|
||||
# Only recalculate RLEs for masks that have changed
|
||||
for i_mask in keep_by_nms:
|
||||
if scores[i_mask] == 0.0:
|
||||
mask_torch = masks[i_mask].unsqueeze(0)
|
||||
rles[i_mask] = mask_to_rle_pytorch(mask_torch)
|
||||
masks = [rle_to_mask(rles[i][0]) for i in keep_by_nms]
|
||||
return masks
|
||||
|
||||
def run_everything_ours(self, img_tensor, model):
|
||||
from segment_anything.utils.amg import mask_to_rle_pytorch
|
||||
img_tensor = img_tensor.squeeze(0)
|
||||
_, original_image_h, original_image_w = img_tensor.shape
|
||||
xy = []
|
||||
for i in range(self.GRID_SIZE):
|
||||
curr_x = 0.5 + i / self.GRID_SIZE * original_image_w
|
||||
for j in range(self.GRID_SIZE):
|
||||
curr_y = 0.5 + j / self.GRID_SIZE * original_image_h
|
||||
xy.append([curr_x, curr_y])
|
||||
|
||||
xy = torch.from_numpy(np.array(xy))
|
||||
points = xy
|
||||
num_pts = xy.shape[0]
|
||||
point_labels = torch.ones(num_pts, 1)
|
||||
with torch.no_grad():
|
||||
predicted_masks, predicted_iou = self.get_predictions_given_embeddings_and_queries(
|
||||
img_tensor,
|
||||
points.reshape(1, num_pts, 1, 2).to(we.device_id),
|
||||
point_labels.reshape(1, num_pts, 1).to(we.device_id),
|
||||
model,
|
||||
)
|
||||
# print('predicted_masks: ', predicted_masks[0][0:1].dtype, predicted_masks[0][0:1].device)
|
||||
rle = [mask_to_rle_pytorch(m[0:1]) for m in predicted_masks]
|
||||
# transform to numpy
|
||||
size, counts = [], []
|
||||
for rle_item in rle:
|
||||
size.append(rle_item[0]['size'])
|
||||
counts += rle_item[0]['counts']
|
||||
counts += '#'
|
||||
predicted_masks = self.process_small_region(rle)
|
||||
return predicted_masks
|
||||
|
||||
def forward(self, image, return_mask=None):
|
||||
return_mask = return_mask if return_mask is not None else self.return_mask
|
||||
if isinstance(image, Image.Image):
|
||||
image = np.array(image)
|
||||
elif isinstance(image, torch.Tensor):
|
||||
image = image.detach().cpu().numpy()
|
||||
elif isinstance(image, np.ndarray):
|
||||
image = image.copy()
|
||||
else:
|
||||
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
|
||||
h, w = image.shape[:2]
|
||||
max_rate = max(float(w) / 1024.0, float(h) / 1024.0)
|
||||
w_ori = int(float(w) / max_rate)
|
||||
h_ori = int(float(h) / max_rate)
|
||||
# image = T.ToTensor()(T.Resize((h_ori, w_ori))(Image.fromarray(image)))
|
||||
# image_pad = T.Pad((0, 0, 1024 - w_ori, 1024 - h_ori))(image)
|
||||
image_pad = T.Pad((0, 0, 1024 - w_ori, 1024 - h_ori))(T.Resize(
|
||||
(h_ori, w_ori))(Image.fromarray(image)))
|
||||
input_image = T.ToTensor()(image_pad)
|
||||
input_image = input_image.unsqueeze(0).to(we.device_id)
|
||||
|
||||
mask_efficient_sam_vits = self.run_everything_ours(
|
||||
input_image, self.efficient_sam_module)
|
||||
annos = []
|
||||
mask_efficient_sam_vits = sorted(list(mask_efficient_sam_vits),
|
||||
key=lambda m: int(m.sum()),
|
||||
reverse=True)
|
||||
mask_efficient_sam_vits = mask_efficient_sam_vits[:256]
|
||||
for mask in mask_efficient_sam_vits:
|
||||
mask_item = mask_utils.encode(
|
||||
np.array(mask[:, :, None], order='F', dtype='uint8'))[0]
|
||||
mask_item['counts'] = mask_item['counts'].decode('utf-8')
|
||||
mask_area = int(mask.sum())
|
||||
annos.append({'mask': mask_item, 'mask_area': mask_area})
|
||||
|
||||
annos = sorted(annos, key=lambda x: x['mask_area'], reverse=True)
|
||||
seg_img = None
|
||||
dominant_palette = []
|
||||
image_pad_np = np.array(image_pad)
|
||||
for idx, anno in enumerate(annos):
|
||||
color = idx
|
||||
if idx > 255:
|
||||
break
|
||||
mask = np.array(mask_utils.decode(anno['mask'])).astype(np.uint8)
|
||||
h, w = mask.shape[:2]
|
||||
if seg_img is None:
|
||||
seg_img = np.ones((h, w, 3)) * 255
|
||||
if self.use_dominant_color:
|
||||
masked_image = cv2.bitwise_and(image_pad_np,
|
||||
image_pad_np,
|
||||
mask=mask)
|
||||
dominant_color = find_dominant_color(masked_image).tolist()
|
||||
dominant_palette.append(dominant_color)
|
||||
seg_img[mask.astype(bool)] = [color, color, color]
|
||||
seg_img = Image.fromarray(seg_img.astype(np.uint8)).convert('L')
|
||||
|
||||
resize_rate = max(float(h_ori) / 1024.0, float(w_ori) / 1024.0)
|
||||
h_new = int(float(h_ori) / resize_rate)
|
||||
w_new = int(float(w_ori) / resize_rate)
|
||||
seg_img = seg_img.crop((0, 0, w_new, h_new))
|
||||
if self.save_mode == 'P':
|
||||
palette = []
|
||||
for i in range(256):
|
||||
if not self.use_dominant_color:
|
||||
palette_item = [random.randint(0, 255) for _ in range(3)]
|
||||
else:
|
||||
palette_item = dominant_palette[i] if i < len(
|
||||
dominant_palette) else [255, 255, 255]
|
||||
palette += palette_item
|
||||
seg_img = seg_img.convert('P')
|
||||
seg_img.putpalette(palette)
|
||||
seg_rgb_img = seg_img.convert('RGB')
|
||||
if return_mask:
|
||||
return {
|
||||
'image': np.array(seg_rgb_img),
|
||||
'mask': np.array(seg_img)
|
||||
}
|
||||
else:
|
||||
return np.array(seg_rgb_img)
|
||||
else:
|
||||
return np.array(seg_img)
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('ANNOTATORS',
|
||||
__class__.__name__,
|
||||
ESAMAnnotator.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
|
||||
|
||||
@ANNOTATORS.register_class()
|
||||
class SAMAnnotatorDraw(BaseAnnotator, metaclass=ABCMeta):
|
||||
para_dict = {}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
from segment_anything import sam_model_registry, SamPredictor
|
||||
from segment_anything.utils.transforms import ResizeLongestSide
|
||||
|
||||
self.transform = ResizeLongestSide(1024)
|
||||
self.task_type = cfg.get('TASK_TYPE', 'input_box')
|
||||
self.sam_model = cfg.get('SAM_MODEL', 'vit_b')
|
||||
pretrained_model = cfg.get('PRETRAINED_MODEL', 'sam_vit_b_01ec64.pth')
|
||||
|
||||
if pretrained_model:
|
||||
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
|
||||
seg_model = sam_model_registry[self.sam_model](checkpoint=local_path).eval().to(we.device_id)
|
||||
self.sam_predictor = SamPredictor(seg_model)
|
||||
|
||||
def forward(self, image, input_box=None, mask=None, task_type=None, multimask_output=False):
|
||||
task_type = task_type if task_type is not None else self.task_type
|
||||
|
||||
if isinstance(image, Image.Image):
|
||||
image = np.array(image)
|
||||
elif isinstance(image, torch.Tensor):
|
||||
image = image.detach().cpu().numpy()
|
||||
elif isinstance(image, np.ndarray):
|
||||
image = image.copy()
|
||||
else:
|
||||
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
|
||||
|
||||
if mask is not None:
|
||||
if isinstance(mask, Image.Image):
|
||||
mask = np.array(mask)
|
||||
elif isinstance(mask, torch.Tensor):
|
||||
mask = mask.detach().cpu().numpy()
|
||||
elif isinstance(mask, np.ndarray):
|
||||
mask = mask.copy()
|
||||
else:
|
||||
raise f'Unsurpport datatype{type(mask)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
|
||||
|
||||
original_size = image.shape[:2]
|
||||
if task_type == 'mask_point':
|
||||
scribble = mask.transpose(2, 1, 0)[0]
|
||||
labeled_array, num_features = ndimage.label(scribble >= 255)
|
||||
centers = ndimage.center_of_mass(scribble, labeled_array, range(1, num_features + 1))
|
||||
point_coords = np.array(centers)
|
||||
point_labels = np.array([1] * len(centers))
|
||||
sample = {'point_coords': point_coords, 'point_labels': point_labels}
|
||||
|
||||
elif task_type == 'mask_box':
|
||||
scribble = mask.transpose(2, 1, 0)[0]
|
||||
labeled_array, num_features = ndimage.label(scribble >= 255)
|
||||
centers = ndimage.center_of_mass(scribble, labeled_array, range(1, num_features + 1))
|
||||
centers = np.array(centers)
|
||||
### (x1, y1, x2, y2)
|
||||
x_min = centers[:, 0].min()
|
||||
x_max = centers[:, 0].max()
|
||||
y_min = centers[:, 1].min()
|
||||
y_max = centers[:, 1].max()
|
||||
bbox = np.array([x_min, y_min, x_max, y_max])
|
||||
sample = {'box': bbox}
|
||||
|
||||
elif task_type == 'input_box':
|
||||
if isinstance(input_box, list):
|
||||
input_box = np.array(input_box)
|
||||
sample = {'box': input_box}
|
||||
|
||||
self.sam_predictor.set_image(image)
|
||||
masks, scores, logits = self.sam_predictor.predict(**sample, multimask_output=True)
|
||||
index = np.argmax(scores)
|
||||
|
||||
ret_data = {
|
||||
"mask": (masks[index]* 255).astype(np.uint8),
|
||||
"score": scores[index]
|
||||
}
|
||||
return ret_data
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('ANNOTATORS',
|
||||
__class__.__name__,
|
||||
SAMAnnotatorDraw.para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,160 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
import math
|
||||
from abc import ABCMeta
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torchvision.transforms as TT
|
||||
from einops import rearrange
|
||||
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||
from scepter.modules.annotator.registry import ANNOTATORS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
|
||||
class SketchNet(nn.Module):
|
||||
def __init__(self, mean, std):
|
||||
assert isinstance(mean, float) and isinstance(std, float)
|
||||
super().__init__()
|
||||
self.mean = mean
|
||||
self.std = std
|
||||
|
||||
# layers
|
||||
self.layers = nn.Sequential(nn.Conv2d(1, 48, 5, 2, 2),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(48, 128, 3, 1, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(128, 128, 3, 1, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(128, 128, 3, 2, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(128, 256, 3, 1, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(256, 256, 3, 1, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(256, 256, 3, 2, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(256, 512, 3, 1, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(512, 1024, 3, 1, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(1024, 1024, 3, 1, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(1024, 1024, 3, 1, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(1024, 1024, 3, 1, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(1024, 512, 3, 1, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(512, 256, 3, 1, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.ConvTranspose2d(256, 256, 4, 2, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(256, 256, 3, 1, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(256, 128, 3, 1, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.ConvTranspose2d(128, 128, 4, 2, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(128, 128, 3, 1, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(128, 48, 3, 1, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.ConvTranspose2d(48, 48, 4, 2, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(48, 24, 3, 1, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(24, 1, 3, 1, 1), nn.Sigmoid())
|
||||
|
||||
def forward(self, x):
|
||||
"""x: [B, 1, H, W] within range [0, 1]. Sketch pixels in dark color.
|
||||
"""
|
||||
x = (x - self.mean) / self.std
|
||||
return self.layers(x)
|
||||
|
||||
|
||||
@ANNOTATORS.register_class()
|
||||
class SketchAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
para_dict = {}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
pretrained_model = cfg.get('PRETRAINED_MODEL', None)
|
||||
self.model = SketchNet(mean=0.9664114577640158,
|
||||
std=0.0858381272736797).eval()
|
||||
if pretrained_model:
|
||||
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
|
||||
state = torch.load(local_path, map_location='cpu')
|
||||
self.model.load_state_dict(state)
|
||||
|
||||
@torch.no_grad()
|
||||
@torch.inference_mode()
|
||||
@torch.autocast('cuda', enabled=False)
|
||||
def forward(self, image):
|
||||
is_batch = False if len(image.shape) == 3 else True
|
||||
if isinstance(image, torch.Tensor):
|
||||
if len(image.shape) == 3:
|
||||
if torch.equal(image[:, :, 0], image[:, :, 1]) and torch.equal(
|
||||
image[:, :, 1], image[:, :, 2]):
|
||||
image = image[:, :, 0]
|
||||
else:
|
||||
raise "Unsurpport input image's shape and each channel is different"
|
||||
elif len(image.shape) == 4:
|
||||
if (torch.equal(image[:, :, :, 0], image[:, :, :, 1])
|
||||
and torch.equal(image[:, :, :, 1], image[:, :, :, 2])):
|
||||
image = image[:, :, :, 0]
|
||||
else:
|
||||
raise "Unsurpport input image's shape and each channel is different"
|
||||
if len(image.shape) == 2:
|
||||
image = rearrange(image, 'h w -> 1 h w')
|
||||
B, H, W = image.shape
|
||||
elif len(image.shape) == 3:
|
||||
B, H, W = image.shape
|
||||
else:
|
||||
raise "Unsurpport input image's shape"
|
||||
elif isinstance(image, np.ndarray):
|
||||
image = image.copy()
|
||||
if len(image.shape) == 3:
|
||||
if np.array_equal(image[:, :, 0],
|
||||
image[:, :, 1]) and np.array_equal(
|
||||
image[:, :, 1], image[:, :, 2]):
|
||||
image = image[:, :, 0]
|
||||
else:
|
||||
raise "Unsurpport input image's shape and each channel is different"
|
||||
elif len(image.shape) == 4:
|
||||
if (np.array_equal(image[:, :, :, 0], image[:, :, :, 1]) and
|
||||
np.array_equal(image[:, :, :, 1], image[:, :, :, 2])):
|
||||
image = image[:, :, :, 0]
|
||||
else:
|
||||
raise "Unsurpport input image's shape and each channel is different"
|
||||
image = torch.from_numpy(image).float()
|
||||
if len(image.shape) == 2:
|
||||
image = rearrange(image, 'h w -> 1 1 h w')
|
||||
elif len(image.shape) == 3:
|
||||
image = rearrange(image, 'b h w -> b 1 h w')
|
||||
else:
|
||||
raise "Unsurpport input image's shape"
|
||||
else:
|
||||
raise "Unsurpport input image's type"
|
||||
image = image.float().div(255)
|
||||
image = image.to(we.device_id)
|
||||
edge = self.model(image)
|
||||
edge = edge.squeeze(dim=1)
|
||||
edge = (edge * 255.0).clip(0, 255)
|
||||
edge = edge.cpu().numpy()
|
||||
edge = edge.astype(np.uint8)
|
||||
if not is_batch:
|
||||
edge = edge.squeeze()
|
||||
return edge[..., None].repeat(3, -1)
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('ANNOTATORS',
|
||||
__class__.__name__,
|
||||
SketchAnnotator.para_dict,
|
||||
set_name=True)
|
||||
Reference in New Issue
Block a user