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)
|
||||
@@ -82,6 +82,8 @@ class BaseDataset(Dataset, metaclass=ABCMeta):
|
||||
overwrite=False)
|
||||
self.worker_id = worker_id
|
||||
self.logger = self.worker_logger
|
||||
self.local_we["seed"] += (worker_id + we.rank)
|
||||
self.seed = self.local_we["seed"]
|
||||
we.set_env(self.local_we)
|
||||
|
||||
@abstractmethod
|
||||
|
||||
@@ -258,6 +258,8 @@ class DataObject(object):
|
||||
if sampler_name == 'MixtureOfSamplers':
|
||||
subsampler_configs = self.data_sampler_config.get(
|
||||
'SUB_SAMPLERS', [])
|
||||
keep_order = self.data_sampler_config.get(
|
||||
'KEEP_ORDER', False)
|
||||
subsamplers = list()
|
||||
subsampler_probs = list()
|
||||
for ssconfig in subsampler_configs:
|
||||
@@ -277,7 +279,7 @@ class DataObject(object):
|
||||
subsampler_probs.append(prob)
|
||||
self.batch_sampler = MixtureOfSamplers(subsamplers,
|
||||
subsampler_probs, rank,
|
||||
seed)
|
||||
seed, keep_order = keep_order)
|
||||
elif sampler_name == 'MultiLevelBatchSampler':
|
||||
self.batch_sampler = self._instantiate_multi_level_batch_sampler(
|
||||
self.data_sampler_config, self.batch_size, rank, seed)
|
||||
|
||||
@@ -523,12 +523,15 @@ class MultiLevelBatchSampler(BaseSampler):
|
||||
class MixtureOfSamplers(BaseSampler):
|
||||
para_dict = {'SUB_SAMPLERS': []}
|
||||
|
||||
def __init__(self, samplers, probabilities, rank=0, seed=8888):
|
||||
def __init__(self, samplers, probabilities, rank=0, seed=8888, keep_order = False):
|
||||
self.samplers = samplers
|
||||
self.iterators = [iter(u) for u in samplers]
|
||||
self.probabilities = probabilities
|
||||
self.seed = seed
|
||||
self.rng = np.random.default_rng(seed + rank)
|
||||
if keep_order:
|
||||
self.rng = np.random.default_rng(seed)
|
||||
else:
|
||||
self.rng = np.random.default_rng(seed + rank)
|
||||
|
||||
def __iter__(self):
|
||||
while True:
|
||||
|
||||
@@ -13,12 +13,6 @@ from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
try:
|
||||
from swift import SwiftModel
|
||||
except Exception:
|
||||
warnings.warn('Import swift failed, please check it.')
|
||||
|
||||
|
||||
class ControlInference():
|
||||
def __init__(self, logger=None):
|
||||
self.logger = logger
|
||||
@@ -26,6 +20,11 @@ class ControlInference():
|
||||
|
||||
# @classmethod
|
||||
def unregister_controllers(self, control_model_ins, diffusion_model):
|
||||
try:
|
||||
from swift import SwiftModel
|
||||
except Exception:
|
||||
warnings.warn('Import swift failed, please check it.')
|
||||
|
||||
self.logger.info('Unloading control model')
|
||||
if isinstance(diffusion_model['model'], SwiftModel):
|
||||
if (hasattr(diffusion_model['model'].base_model, 'control_blocks')
|
||||
@@ -42,6 +41,11 @@ class ControlInference():
|
||||
|
||||
# @classmethod
|
||||
def register_controllers(self, control_model_ins, diffusion_model):
|
||||
try:
|
||||
from swift import SwiftModel
|
||||
except Exception:
|
||||
warnings.warn('Import swift failed, please check it.')
|
||||
|
||||
self.logger.info('Loading control model')
|
||||
if control_model_ins is None or control_model_ins == '':
|
||||
self.unregister_controllers(control_model_ins, diffusion_model)
|
||||
|
||||
@@ -84,14 +84,13 @@ class DiffusionInference():
|
||||
|
||||
def redefine_paras(self, cfg):
|
||||
if cfg.get('PRETRAINED_MODEL', None):
|
||||
assert FS.isfile(cfg.PRETRAINED_MODEL)
|
||||
with FS.get_from(cfg.PRETRAINED_MODEL,
|
||||
wait_finish=True) as local_path:
|
||||
if local_path.endswith('safetensors'):
|
||||
from safetensors.torch import load_file as load_safetensors
|
||||
sd = load_safetensors(local_path)
|
||||
else:
|
||||
sd = torch.load(local_path, map_location='cpu')
|
||||
sd = torch.load(local_path, map_location='cpu', weights_only=True)
|
||||
first_stage_model_path = os.path.join(
|
||||
os.path.dirname(local_path), 'first_stage_model.pth')
|
||||
cond_stage_model_path = os.path.join(
|
||||
@@ -311,7 +310,7 @@ class DiffusionInference():
|
||||
module_paras = {}
|
||||
if cfg is not None:
|
||||
self.paras = cfg.PARAS
|
||||
self.input = {k.lower(): v for k, v in cfg.INPUT.items()}
|
||||
self.input = {k.lower(): dict(v).get('DEFAULT', None) if isinstance(v, (dict, OrderedDict)) else v for k, v in cfg.INPUT.items()}
|
||||
self.output = {k.lower(): v for k, v in cfg.OUTPUT.items()}
|
||||
module_paras = cfg.MODULES_PARAS
|
||||
return module_paras
|
||||
|
||||
@@ -0,0 +1,216 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import math
|
||||
import random
|
||||
import torch
|
||||
from scepter.modules.utils.distribute import we
|
||||
from .control_inference import ControlInference
|
||||
from .diffusion_inference import DiffusionInference, get_model
|
||||
from .tuner_inference import TunerInference
|
||||
from scepter.modules.model.registry import DIFFUSIONS, TOKENIZERS
|
||||
|
||||
|
||||
class FluxInference(DiffusionInference):
|
||||
def __init__(self, logger=None):
|
||||
self.logger = logger
|
||||
self.is_redefine_paras = False
|
||||
self.loaded_model = {}
|
||||
self.loaded_model_name = [
|
||||
'diffusion_model', 'first_stage_model', 'cond_stage_model'
|
||||
]
|
||||
self.tuner_infer = TunerInference(self.logger)
|
||||
self.control_infer = ControlInference(self.logger)
|
||||
|
||||
def init_from_cfg(self, cfg):
|
||||
self.name = cfg.NAME
|
||||
self.is_default = cfg.get('IS_DEFAULT', False)
|
||||
module_paras = self.load_default(cfg.get('DEFAULT_PARAS', None))
|
||||
assert cfg.have('MODEL')
|
||||
if self.is_redefine_paras:
|
||||
cfg.MODEL = self.redefine_paras(cfg.MODEL)
|
||||
self.diffusion_model = self.infer_model(
|
||||
cfg.MODEL.DIFFUSION_MODEL, module_paras.get(
|
||||
'DIFFUSION_MODEL',
|
||||
None)) if cfg.MODEL.have('DIFFUSION_MODEL') else None
|
||||
self.first_stage_model = self.infer_model(
|
||||
cfg.MODEL.FIRST_STAGE_MODEL,
|
||||
module_paras.get(
|
||||
'FIRST_STAGE_MODEL',
|
||||
None)) if cfg.MODEL.have('FIRST_STAGE_MODEL') else None
|
||||
self.cond_stage_model = self.infer_model(
|
||||
cfg.MODEL.COND_STAGE_MODEL,
|
||||
module_paras.get(
|
||||
'COND_STAGE_MODEL',
|
||||
None)) if cfg.MODEL.have('COND_STAGE_MODEL') else None
|
||||
self.refiner_cond_model = self.infer_model(
|
||||
cfg.MODEL.REFINER_COND_MODEL,
|
||||
module_paras.get(
|
||||
'REFINER_COND_MODEL',
|
||||
None)) if cfg.MODEL.have('REFINER_COND_MODEL') else None
|
||||
self.refiner_diffusion_model = self.infer_model(
|
||||
cfg.MODEL.REFINER_MODEL, module_paras.get(
|
||||
'REFINER_MODEL',
|
||||
None)) if cfg.MODEL.have('REFINER_MODEL') else None
|
||||
self.tokenizer = TOKENIZERS.build(
|
||||
cfg.MODEL.TOKENIZER,
|
||||
logger=self.logger) if cfg.MODEL.have('TOKENIZER') else None
|
||||
|
||||
if self.tokenizer is not None:
|
||||
self.cond_stage_model['cfg'].KWARGS = {
|
||||
'vocab_size': self.tokenizer.vocab_size
|
||||
}
|
||||
self.diffusion = DIFFUSIONS.build(cfg.MODEL.DIFFUSION, logger=self.logger) \
|
||||
if cfg.MODEL.have('DIFFUSION') else None
|
||||
assert self.diffusion is not None
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_first_stage(self, x, **kwargs):
|
||||
_, dtype = self.get_function_info(self.first_stage_model, 'encode')
|
||||
with torch.autocast('cuda',
|
||||
enabled= dtype in ('float16', 'bfloat16'),
|
||||
dtype=getattr(torch, dtype)):
|
||||
z = get_model(self.first_stage_model).encode(x)
|
||||
if isinstance(z, (tuple, list)):
|
||||
z = z[0]
|
||||
return z
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_first_stage(self, z):
|
||||
_, dtype = self.get_function_info(self.first_stage_model, 'decode')
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype in ('float16', 'bfloat16'),
|
||||
dtype=getattr(torch, dtype)):
|
||||
return get_model(self.first_stage_model).decode(z)
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(self,
|
||||
input,
|
||||
num_samples=1,
|
||||
cat_uc=True,
|
||||
tuner_model=None,
|
||||
control_model=None,
|
||||
**kwargs):
|
||||
|
||||
value_input = copy.deepcopy(self.input)
|
||||
value_input.update(input)
|
||||
print(value_input)
|
||||
height, width = value_input['target_size_as_tuple']
|
||||
value_output = copy.deepcopy(self.output)
|
||||
# register tuner
|
||||
if tuner_model is not None and tuner_model != '' and len(
|
||||
tuner_model) > 0:
|
||||
if not isinstance(tuner_model, list):
|
||||
tuner_model = [tuner_model]
|
||||
self.dynamic_load(self.diffusion_model, 'diffusion_model')
|
||||
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
|
||||
self.tuner_infer.register_tuner(tuner_model, self.diffusion_model,
|
||||
self.cond_stage_model)
|
||||
self.dynamic_unload(self.diffusion_model,
|
||||
'diffusion_model',
|
||||
skip_loaded=True)
|
||||
self.dynamic_unload(self.cond_stage_model,
|
||||
'cond_stage_model',
|
||||
skip_loaded=True)
|
||||
|
||||
# cond stage
|
||||
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
|
||||
function_name, dtype = self.get_function_info(self.cond_stage_model)
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype == 'float16',
|
||||
dtype=getattr(torch, dtype)):
|
||||
ctx = getattr(get_model(self.cond_stage_model),
|
||||
function_name)(value_input['prompt'])
|
||||
self.dynamic_unload(self.cond_stage_model,
|
||||
'cond_stage_model',
|
||||
skip_loaded=True)
|
||||
|
||||
# get noise
|
||||
seed = kwargs.pop('seed', -1)
|
||||
g = torch.Generator(device=we.device_id)
|
||||
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
||||
g.manual_seed(seed)
|
||||
if 'seed' in value_output:
|
||||
value_output['seed'] = seed
|
||||
for sample_id in range(num_samples):
|
||||
if self.diffusion_model is not None:
|
||||
noise = torch.randn(
|
||||
num_samples,
|
||||
16,
|
||||
# allow for packing
|
||||
2 * math.ceil(height / 16),
|
||||
2 * math.ceil(width / 16),
|
||||
device=we.device_id,
|
||||
dtype=getattr(torch, dtype),
|
||||
generator=torch.Generator(device=we.device_id).manual_seed(seed),
|
||||
)
|
||||
self.dynamic_load(self.diffusion_model, 'diffusion_model')
|
||||
# UNet use input n_prompt
|
||||
function_name, dtype = self.get_function_info(
|
||||
self.diffusion_model)
|
||||
with torch.autocast('cuda',
|
||||
enabled= dtype in ('float16', 'bfloat16'),
|
||||
dtype=getattr(torch, dtype)):
|
||||
solver_sample = value_input.get('sample', 'flow_eluer')
|
||||
sample_steps = value_input.get('sample_steps', 20)
|
||||
guide_scale = value_input.get('guide_scale', 3.5)
|
||||
if guide_scale is not None:
|
||||
guide_scale = torch.full((noise.shape[0],), guide_scale, device=noise.device,
|
||||
dtype=noise.dtype)
|
||||
else:
|
||||
guide_scale = None
|
||||
|
||||
latent = self.diffusion.sample(
|
||||
noise=noise,
|
||||
sampler=solver_sample,
|
||||
model=get_model(self.diffusion_model),
|
||||
model_kwargs={"cond": ctx, "guidance": guide_scale},
|
||||
steps=sample_steps,
|
||||
show_progress=True,
|
||||
guide_scale=guide_scale,
|
||||
return_intermediate=None,
|
||||
**kwargs)
|
||||
|
||||
self.dynamic_unload(self.diffusion_model,
|
||||
'diffusion_model',
|
||||
skip_loaded=True)
|
||||
|
||||
if 'latent' in value_output:
|
||||
if value_output['latent'] is None or (
|
||||
isinstance(value_output['latent'], list)
|
||||
and len(value_output['latent']) < 1):
|
||||
value_output['latent'] = []
|
||||
value_output['latent'].append(latent)
|
||||
|
||||
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
x_samples = self.decode_first_stage(latent).float()
|
||||
self.dynamic_unload(self.first_stage_model,
|
||||
'first_stage_model',
|
||||
skip_loaded=True)
|
||||
images = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
if 'images' in value_output:
|
||||
if value_output['images'] is None or (
|
||||
isinstance(value_output['images'], list)
|
||||
and len(value_output['images']) < 1):
|
||||
value_output['images'] = []
|
||||
value_output['images'].append(images)
|
||||
|
||||
for k, v in value_output.items():
|
||||
if isinstance(v, list):
|
||||
value_output[k] = torch.cat(v, dim=0)
|
||||
if isinstance(v, torch.Tensor):
|
||||
value_output[k] = v.cpu()
|
||||
|
||||
# unregister tuner
|
||||
if tuner_model is not None and tuner_model != '' and len(
|
||||
tuner_model) > 0:
|
||||
self.tuner_infer.unregister_tuner(tuner_model,
|
||||
self.diffusion_model,
|
||||
self.cond_stage_model)
|
||||
|
||||
# unregister control
|
||||
if control_model is not None and control_model != '':
|
||||
self.control_infer.unregister_controllers(control_model,
|
||||
self.diffusion_model)
|
||||
|
||||
return value_output
|
||||
@@ -40,7 +40,7 @@ class LargenInference(DiffusionInference):
|
||||
from safetensors.torch import load_file as load_safetensors
|
||||
sd = load_safetensors(local_path)
|
||||
else:
|
||||
sd = torch.load(local_path, map_location='cpu')
|
||||
sd = torch.load(local_path, map_location='cpu', weights_only=True)
|
||||
|
||||
if 'model' in sd:
|
||||
sd = sd['model']
|
||||
|
||||
@@ -84,9 +84,11 @@ class PixArtInference(DiffusionInference):
|
||||
function_name)(value_input['prompt'],
|
||||
return_mask=True)
|
||||
context['crossattn'] = cont.float()
|
||||
self.dynamic_load(self.diffusion_model, 'diffusion_model')
|
||||
null_context['crossattn'] = get_model(
|
||||
self.diffusion_model).y_embedder.y_embedding[None].repeat(
|
||||
num_samples, 1, 1)
|
||||
self.dynamic_unload(self.diffusion_model, 'diffusion_model')
|
||||
self.dynamic_unload(self.cond_stage_model,
|
||||
'cond_stage_model',
|
||||
skip_loaded=True)
|
||||
|
||||
@@ -13,10 +13,6 @@ try:
|
||||
from peft.utils import CONFIG_NAME, SAFETENSORS_WEIGHTS_NAME, WEIGHTS_NAME
|
||||
except Exception as e:
|
||||
warnings.warn(f'Import peft error, please deal with this problem: {e}')
|
||||
try:
|
||||
from swift import Swift, SwiftModel
|
||||
except Exception as e:
|
||||
warnings.warn(f'Import swift error, please deal with this problem: {e}')
|
||||
|
||||
|
||||
class TunerInference():
|
||||
@@ -27,6 +23,11 @@ class TunerInference():
|
||||
# @classmethod
|
||||
def unregister_tuner(self, tuner_model_list, diffusion_model,
|
||||
cond_stage_model):
|
||||
try:
|
||||
from swift import SwiftModel
|
||||
except Exception as e:
|
||||
warnings.warn(f'Import swift error, please deal with this problem: {e}')
|
||||
|
||||
self.logger.info('Unloading tuner model')
|
||||
if isinstance(diffusion_model['model'], SwiftModel):
|
||||
for adapter_name in diffusion_model['model'].adapters:
|
||||
@@ -41,6 +42,11 @@ class TunerInference():
|
||||
# @classmethod
|
||||
def register_tuner(self, tuner_model_list, diffusion_model,
|
||||
cond_stage_model):
|
||||
try:
|
||||
from swift import Swift
|
||||
except Exception as e:
|
||||
warnings.warn(f'Import swift error, please deal with this problem: {e}')
|
||||
|
||||
self.logger.info('Loading tuner model')
|
||||
if len(tuner_model_list) < 1:
|
||||
self.unregister_tuner(tuner_model_list, diffusion_model,
|
||||
@@ -137,7 +143,7 @@ class TunerInference():
|
||||
state_dict = {}
|
||||
is_bin_file = True
|
||||
if os.path.isfile(bin_file):
|
||||
state_dict = torch.load(bin_file)
|
||||
state_dict = torch.load(bin_file, weights_only=True)
|
||||
elif os.path.isfile(safe_file):
|
||||
is_bin_file = False
|
||||
from safetensors.torch import \
|
||||
|
||||
@@ -2,4 +2,4 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
from scepter.modules.model import (backbone, embedder, head, loss, metric,
|
||||
neck, network, tokenizer, tuner)
|
||||
neck, network, tokenizer, tuner, diffusion)
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.backbone import (autoencoder, image, mmdit, pixart,
|
||||
unet, utils, video)
|
||||
unet, utils, video, flux)
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
from .flux import Flux
|
||||
@@ -0,0 +1,245 @@
|
||||
import math
|
||||
from functools import partial
|
||||
|
||||
import torch
|
||||
from einops import rearrange, repeat
|
||||
from scepter.modules.model.base_model import BaseModel
|
||||
from scepter.modules.model.registry import BACKBONES
|
||||
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 torch import Tensor, nn
|
||||
from torch.utils.checkpoint import checkpoint_sequential
|
||||
|
||||
from .layers import (DoubleStreamBlock, EmbedND, LastLayer,
|
||||
MLPEmbedder, SingleStreamBlock,
|
||||
timestep_embedding)
|
||||
|
||||
@BACKBONES.register_class()
|
||||
class Flux(BaseModel):
|
||||
"""
|
||||
Transformer backbone Diffusion model with RoPE.
|
||||
"""
|
||||
para_dict = {
|
||||
"IN_CHANNELS": {
|
||||
"value": 64,
|
||||
"description": "model's input channels."
|
||||
},
|
||||
"OUT_CHANNELS": {
|
||||
"value": 64,
|
||||
"description": "model's output channels."
|
||||
},
|
||||
"HIDDEN_SIZE": {
|
||||
"value": 1024,
|
||||
"description": "model's hidden size."
|
||||
},
|
||||
"NUM_HEADS": {
|
||||
"value": 16,
|
||||
"description": "number of heads in the transformer."
|
||||
},
|
||||
"AXES_DIM": {
|
||||
"value": [16, 56, 56],
|
||||
"description": "dimensions of the axes of the positional encoding."
|
||||
},
|
||||
"THETA": {
|
||||
"value": 10_000,
|
||||
"description": "theta for positional encoding."
|
||||
},
|
||||
"VEC_IN_DIM": {
|
||||
"value": 768,
|
||||
"description": "dimension of the vector input."
|
||||
},
|
||||
"GUIDANCE_EMBED": {
|
||||
"value": False,
|
||||
"description": "whether to use guidance embedding."
|
||||
},
|
||||
"CONTEXT_IN_DIM": {
|
||||
"value": 4096,
|
||||
"description": "dimension of the context input."
|
||||
},
|
||||
"MLP_RATIO": {
|
||||
"value": 4.0,
|
||||
"description": "ratio of mlp hidden size to hidden size."
|
||||
},
|
||||
"QKV_BIAS": {
|
||||
"value": True,
|
||||
"description": "whether to use bias in qkv projection."
|
||||
},
|
||||
"DEPTH": {
|
||||
"value": 19,
|
||||
"description": "number of transformer blocks."
|
||||
},
|
||||
"DEPTH_SINGLE_BLOCKS": {
|
||||
"value": 38,
|
||||
"description": "number of transformer blocks in the single stream block."
|
||||
},
|
||||
"USE_GRAD_CHECKPOINT": {
|
||||
"value": False,
|
||||
"description": "whether to use gradient checkpointing."
|
||||
}
|
||||
}
|
||||
def __init__(
|
||||
self,
|
||||
cfg,
|
||||
logger = None
|
||||
):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.in_channels = cfg.IN_CHANNELS
|
||||
self.out_channels = cfg.get("OUT_CHANNELS", self.in_channels)
|
||||
hidden_size = cfg.get("HIDDEN_SIZE", 1024)
|
||||
num_heads = cfg.get("NUM_HEADS", 16)
|
||||
axes_dim = cfg.AXES_DIM
|
||||
theta = cfg.THETA
|
||||
vec_in_dim = cfg.VEC_IN_DIM
|
||||
self.guidance_embed = cfg.GUIDANCE_EMBED
|
||||
context_in_dim = cfg.CONTEXT_IN_DIM
|
||||
mlp_ratio = cfg.MLP_RATIO
|
||||
qkv_bias = cfg.QKV_BIAS
|
||||
depth = cfg.DEPTH
|
||||
depth_single_blocks = cfg.DEPTH_SINGLE_BLOCKS
|
||||
self.use_grad_checkpoint = cfg.get("USE_GRAD_CHECKPOINT", False)
|
||||
|
||||
if hidden_size % num_heads != 0:
|
||||
raise ValueError(
|
||||
f"Hidden size {hidden_size} must be divisible by num_heads {num_heads}"
|
||||
)
|
||||
pe_dim = hidden_size // num_heads
|
||||
if sum(axes_dim) != pe_dim:
|
||||
raise ValueError(f"Got {axes_dim} but expected positional dim {pe_dim}")
|
||||
self.hidden_size = hidden_size
|
||||
self.num_heads = num_heads
|
||||
self.pe_embedder = EmbedND(dim=pe_dim, theta=theta, axes_dim= axes_dim)
|
||||
self.img_in = nn.Linear(self.in_channels, self.hidden_size, bias=True)
|
||||
self.time_in = MLPEmbedder(in_dim=256, hidden_dim=self.hidden_size)
|
||||
self.vector_in = MLPEmbedder(vec_in_dim, self.hidden_size)
|
||||
self.guidance_in = (
|
||||
MLPEmbedder(in_dim=256, hidden_dim=self.hidden_size) if self.guidance_embed else nn.Identity()
|
||||
)
|
||||
self.txt_in = nn.Linear(context_in_dim, self.hidden_size)
|
||||
|
||||
self.double_blocks = nn.ModuleList(
|
||||
[
|
||||
DoubleStreamBlock(
|
||||
self.hidden_size,
|
||||
self.num_heads,
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
)
|
||||
for _ in range(depth)
|
||||
]
|
||||
)
|
||||
|
||||
self.single_blocks = nn.ModuleList(
|
||||
[
|
||||
SingleStreamBlock(self.hidden_size, self.num_heads, mlp_ratio=mlp_ratio)
|
||||
for _ in range(depth_single_blocks)
|
||||
]
|
||||
)
|
||||
|
||||
self.final_layer = LastLayer(self.hidden_size, 1, self.out_channels)
|
||||
|
||||
def prepare_input(self, x, context, y, x_shape=None):
|
||||
# x.shape [6, 16, 16, 16] target is [6, 16, 768, 1360]
|
||||
bs, c, h, w = x.shape
|
||||
x = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=2, pw=2)
|
||||
x_id = torch.zeros(h // 2, w // 2, 3)
|
||||
x_id[..., 1] = x_id[..., 1] + torch.arange(h // 2)[:, None]
|
||||
x_id[..., 2] = x_id[..., 2] + torch.arange(w // 2)[None, :]
|
||||
x_ids = repeat(x_id, "h w c -> b (h w) c", b=bs)
|
||||
txt_ids = torch.zeros(bs, context.shape[1], 3)
|
||||
return x, x_ids.to(x), context.to(x), txt_ids.to(x), y.to(x), h, w
|
||||
|
||||
def unpack(self, x: Tensor, height: int, width: int) -> Tensor:
|
||||
return rearrange(
|
||||
x,
|
||||
"b (h w) (c ph pw) -> b c (h ph) (w pw)",
|
||||
h=math.ceil(height/2),
|
||||
w=math.ceil(width/2),
|
||||
ph=2,
|
||||
pw=2,
|
||||
)
|
||||
|
||||
def load_pretrained_model(self, pretrained_model):
|
||||
if next(self.parameters()).device.type == 'meta':
|
||||
map_location = we.device_id
|
||||
else:
|
||||
map_location = "cpu"
|
||||
if pretrained_model is not None:
|
||||
with FS.get_from(pretrained_model, wait_finish=True) as local_model:
|
||||
if local_model.endswith('safetensors'):
|
||||
from safetensors.torch import load_file as load_safetensors
|
||||
sd = load_safetensors(local_model, device=map_location)
|
||||
else:
|
||||
sd = torch.load(local_model, map_location=map_location)
|
||||
missing, unexpected = self.load_state_dict(sd, strict=False, assign=True)
|
||||
self.logger.info(
|
||||
f'Restored from {pretrained_model} with {len(missing)} missing and {len(unexpected)} unexpected keys'
|
||||
)
|
||||
if len(missing) > 0:
|
||||
self.logger.info(f'Missing Keys:\n {missing}')
|
||||
if len(unexpected) > 0:
|
||||
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Tensor,
|
||||
t: Tensor,
|
||||
cond: dict = {},
|
||||
guidance: Tensor | None = None,
|
||||
gc_seg: int = 0
|
||||
) -> Tensor:
|
||||
x, x_ids, txt, txt_ids, y, h, w = self.prepare_input(x, cond["context"], cond["y"])
|
||||
# running on sequences img
|
||||
x = self.img_in(x)
|
||||
vec = self.time_in(timestep_embedding(t, 256))
|
||||
if self.guidance_embed:
|
||||
if guidance is None:
|
||||
raise ValueError("Didn't get guidance strength for guidance distilled model.")
|
||||
vec = vec + self.guidance_in(timestep_embedding(guidance, 256))
|
||||
vec = vec + self.vector_in(y)
|
||||
txt = self.txt_in(txt)
|
||||
ids = torch.cat((txt_ids, x_ids), dim=1)
|
||||
pe = self.pe_embedder(ids)
|
||||
kwargs = dict(
|
||||
vec=vec,
|
||||
pe=pe,
|
||||
txt_length=txt.shape[1],
|
||||
)
|
||||
x = torch.cat((txt, x), 1)
|
||||
if self.use_grad_checkpoint and gc_seg >= 0:
|
||||
x = checkpoint_sequential(
|
||||
functions=[partial(block, **kwargs) for block in self.double_blocks],
|
||||
segments=gc_seg if gc_seg > 0 else len(self.double_blocks),
|
||||
input=x,
|
||||
use_reentrant=False
|
||||
)
|
||||
else:
|
||||
for block in self.double_blocks:
|
||||
x = block(x, **kwargs)
|
||||
|
||||
kwargs = dict(
|
||||
vec=vec,
|
||||
pe=pe,
|
||||
)
|
||||
|
||||
if self.use_grad_checkpoint and gc_seg >= 0:
|
||||
x = checkpoint_sequential(
|
||||
functions=[partial(block, **kwargs) for block in self.single_blocks],
|
||||
segments=gc_seg if gc_seg > 0 else len(self.single_blocks),
|
||||
input=x,
|
||||
use_reentrant=False
|
||||
)
|
||||
else:
|
||||
for block in self.single_blocks:
|
||||
x = block(x, **kwargs)
|
||||
x = x[:, txt.shape[1] :, ...]
|
||||
x = self.final_layer(x, vec) # (N, T, patch_size ** 2 * out_channels) 6 64 64
|
||||
x = self.unpack(x, h, w)
|
||||
return x
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('MODEL',
|
||||
__class__.__name__,
|
||||
Flux.para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,282 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from torch import Tensor, nn
|
||||
import torch
|
||||
from einops import rearrange, repeat
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
def attention(q: Tensor, k: Tensor, v: Tensor, pe: Tensor, mask: Tensor | None = None) -> Tensor:
|
||||
q, k = apply_rope(q, k, pe)
|
||||
x = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask)
|
||||
x = torch.nan_to_num(x, nan=0.0, posinf=1e10, neginf=-1e10)
|
||||
x = rearrange(x, "B H L D -> B L (H D)")
|
||||
return x
|
||||
|
||||
|
||||
def rope(pos: Tensor, dim: int, theta: int) -> Tensor:
|
||||
assert dim % 2 == 0
|
||||
scale = torch.arange(0, dim, 2, dtype=torch.float64, device=pos.device) / dim
|
||||
omega = 1.0 / (theta**scale)
|
||||
out = torch.einsum("...n,d->...nd", pos, omega)
|
||||
out = torch.stack([torch.cos(out), -torch.sin(out), torch.sin(out), torch.cos(out)], dim=-1)
|
||||
out = rearrange(out, "b n d (i j) -> b n d i j", i=2, j=2)
|
||||
return out.float()
|
||||
|
||||
|
||||
def apply_rope(xq: Tensor, xk: Tensor, freqs_cis: Tensor) -> tuple[Tensor, Tensor]:
|
||||
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)
|
||||
xk_ = xk.float().reshape(*xk.shape[:-1], -1, 1, 2)
|
||||
xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1]
|
||||
xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1]
|
||||
return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(*xk.shape).type_as(xk)
|
||||
|
||||
class EmbedND(nn.Module):
|
||||
def __init__(self, dim: int, theta: int, axes_dim: list[int]):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.theta = theta
|
||||
self.axes_dim = axes_dim
|
||||
|
||||
def forward(self, ids: Tensor) -> Tensor:
|
||||
n_axes = ids.shape[-1]
|
||||
emb = torch.cat(
|
||||
[rope(ids[..., i], self.axes_dim[i], self.theta) for i in range(n_axes)],
|
||||
dim=-3,
|
||||
)
|
||||
|
||||
return emb.unsqueeze(1)
|
||||
|
||||
|
||||
def timestep_embedding(t: Tensor, dim, max_period=10000, time_factor: float = 1000.0):
|
||||
"""
|
||||
Create sinusoidal timestep embeddings.
|
||||
:param t: a 1-D Tensor of N indices, one per batch element.
|
||||
These may be fractional.
|
||||
:param dim: the dimension of the output.
|
||||
:param max_period: controls the minimum frequency of the embeddings.
|
||||
:return: an (N, D) Tensor of positional embeddings.
|
||||
"""
|
||||
t = time_factor * t
|
||||
half = dim // 2
|
||||
freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half).to(
|
||||
t.device
|
||||
)
|
||||
|
||||
args = t[:, None].float() * freqs[None]
|
||||
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
||||
if dim % 2:
|
||||
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
||||
if torch.is_floating_point(t):
|
||||
embedding = embedding.to(t)
|
||||
return embedding
|
||||
|
||||
|
||||
class MLPEmbedder(nn.Module):
|
||||
def __init__(self, in_dim: int, hidden_dim: int):
|
||||
super().__init__()
|
||||
self.in_layer = nn.Linear(in_dim, hidden_dim, bias=True)
|
||||
self.silu = nn.SiLU()
|
||||
self.out_layer = nn.Linear(hidden_dim, hidden_dim, bias=True)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
return self.out_layer(self.silu(self.in_layer(x)))
|
||||
|
||||
|
||||
class RMSNorm(torch.nn.Module):
|
||||
def __init__(self, dim: int):
|
||||
super().__init__()
|
||||
self.scale = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def forward(self, x: Tensor):
|
||||
x_dtype = x.dtype
|
||||
x = x.float()
|
||||
rrms = torch.rsqrt(torch.mean(x**2, dim=-1, keepdim=True) + 1e-6)
|
||||
return (x * rrms).to(dtype=x_dtype) * self.scale
|
||||
|
||||
|
||||
class QKNorm(torch.nn.Module):
|
||||
def __init__(self, dim: int):
|
||||
super().__init__()
|
||||
self.query_norm = RMSNorm(dim)
|
||||
self.key_norm = RMSNorm(dim)
|
||||
|
||||
def forward(self, q: Tensor, k: Tensor, v: Tensor) -> tuple[Tensor, Tensor]:
|
||||
q = self.query_norm(q)
|
||||
k = self.key_norm(k)
|
||||
return q.to(v), k.to(v)
|
||||
|
||||
|
||||
class SelfAttention(nn.Module):
|
||||
def __init__(self, dim: int, num_heads: int = 8, qkv_bias: bool = False):
|
||||
super().__init__()
|
||||
self.num_heads = num_heads
|
||||
head_dim = dim // num_heads
|
||||
|
||||
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
|
||||
self.norm = QKNorm(head_dim)
|
||||
self.proj = nn.Linear(dim, dim)
|
||||
|
||||
def forward(self, x: Tensor, pe: Tensor, mask: Tensor | None = None) -> Tensor:
|
||||
qkv = self.qkv(x)
|
||||
q, k, v = rearrange(qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
|
||||
q, k = self.norm(q, k, v)
|
||||
x = attention(q, k, v, pe=pe, mask=mask)
|
||||
x = self.proj(x)
|
||||
return x
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModulationOut:
|
||||
shift: Tensor
|
||||
scale: Tensor
|
||||
gate: Tensor
|
||||
|
||||
|
||||
class Modulation(nn.Module):
|
||||
def __init__(self, dim: int, double: bool):
|
||||
super().__init__()
|
||||
self.is_double = double
|
||||
self.multiplier = 6 if double else 3
|
||||
self.lin = nn.Linear(dim, self.multiplier * dim, bias=True)
|
||||
|
||||
def forward(self, vec: Tensor) -> tuple[ModulationOut, ModulationOut | None]:
|
||||
out = self.lin(nn.functional.silu(vec))[:, None, :].chunk(self.multiplier, dim=-1)
|
||||
|
||||
return (
|
||||
ModulationOut(*out[:3]),
|
||||
ModulationOut(*out[3:]) if self.is_double else None,
|
||||
)
|
||||
|
||||
|
||||
class DoubleStreamBlock(nn.Module):
|
||||
def __init__(self, hidden_size: int, num_heads: int, mlp_ratio: float, qkv_bias: bool = False):
|
||||
super().__init__()
|
||||
|
||||
mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
||||
self.num_heads = num_heads
|
||||
self.hidden_size = hidden_size
|
||||
self.img_mod = Modulation(hidden_size, double=True)
|
||||
self.img_norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.img_attn = SelfAttention(dim=hidden_size, num_heads=num_heads, qkv_bias=qkv_bias)
|
||||
|
||||
self.img_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.img_mlp = nn.Sequential(
|
||||
nn.Linear(hidden_size, mlp_hidden_dim, bias=True),
|
||||
nn.GELU(approximate="tanh"),
|
||||
nn.Linear(mlp_hidden_dim, hidden_size, bias=True),
|
||||
)
|
||||
|
||||
self.txt_mod = Modulation(hidden_size, double=True)
|
||||
self.txt_norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.txt_attn = SelfAttention(dim=hidden_size, num_heads=num_heads, qkv_bias=qkv_bias)
|
||||
|
||||
self.txt_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.txt_mlp = nn.Sequential(
|
||||
nn.Linear(hidden_size, mlp_hidden_dim, bias=True),
|
||||
nn.GELU(approximate="tanh"),
|
||||
nn.Linear(mlp_hidden_dim, hidden_size, bias=True),
|
||||
)
|
||||
|
||||
def forward(self, x: Tensor, vec: Tensor, pe: Tensor, mask: Tensor = None, txt_length = None):
|
||||
img_mod1, img_mod2 = self.img_mod(vec)
|
||||
txt_mod1, txt_mod2 = self.txt_mod(vec)
|
||||
|
||||
txt, img = x[:, :txt_length], x[:, txt_length:]
|
||||
|
||||
# prepare image for attention
|
||||
img_modulated = self.img_norm1(img)
|
||||
img_modulated = (1 + img_mod1.scale) * img_modulated + img_mod1.shift
|
||||
img_qkv = self.img_attn.qkv(img_modulated)
|
||||
img_q, img_k, img_v = rearrange(img_qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
|
||||
img_q, img_k = self.img_attn.norm(img_q, img_k, img_v)
|
||||
# prepare txt for attention
|
||||
txt_modulated = self.txt_norm1(txt)
|
||||
txt_modulated = (1 + txt_mod1.scale) * txt_modulated + txt_mod1.shift
|
||||
txt_qkv = self.txt_attn.qkv(txt_modulated)
|
||||
txt_q, txt_k, txt_v = rearrange(txt_qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
|
||||
txt_q, txt_k = self.txt_attn.norm(txt_q, txt_k, txt_v)
|
||||
|
||||
# run actual attention
|
||||
q = torch.cat((txt_q, img_q), dim=2)
|
||||
k = torch.cat((txt_k, img_k), dim=2)
|
||||
v = torch.cat((txt_v, img_v), dim=2)
|
||||
if mask is not None:
|
||||
mask = repeat(mask, 'B L S-> B H L S', H=self.num_heads)
|
||||
attn = attention(q, k, v, pe=pe, mask = mask)
|
||||
txt_attn, img_attn = attn[:, : txt.shape[1]], attn[:, txt.shape[1] :]
|
||||
|
||||
# calculate the img bloks
|
||||
img = img + img_mod1.gate * self.img_attn.proj(img_attn)
|
||||
img = img + img_mod2.gate * self.img_mlp((1 + img_mod2.scale) * self.img_norm2(img) + img_mod2.shift)
|
||||
|
||||
# calculate the txt bloks
|
||||
txt = txt + txt_mod1.gate * self.txt_attn.proj(txt_attn)
|
||||
txt = txt + txt_mod2.gate * self.txt_mlp((1 + txt_mod2.scale) * self.txt_norm2(txt) + txt_mod2.shift)
|
||||
x = torch.cat((txt, img), 1)
|
||||
return x
|
||||
|
||||
|
||||
class SingleStreamBlock(nn.Module):
|
||||
"""
|
||||
A DiT block with parallel linear layers as described in
|
||||
https://arxiv.org/abs/2302.05442 and adapted modulation interface.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
num_heads: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
qk_scale: float | None = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.hidden_dim = hidden_size
|
||||
self.num_heads = num_heads
|
||||
head_dim = hidden_size // num_heads
|
||||
self.scale = qk_scale or head_dim**-0.5
|
||||
|
||||
self.mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
||||
# qkv and mlp_in
|
||||
self.linear1 = nn.Linear(hidden_size, hidden_size * 3 + self.mlp_hidden_dim)
|
||||
# proj and mlp_out
|
||||
self.linear2 = nn.Linear(hidden_size + self.mlp_hidden_dim, hidden_size)
|
||||
|
||||
self.norm = QKNorm(head_dim)
|
||||
|
||||
self.hidden_size = hidden_size
|
||||
self.pre_norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
|
||||
self.mlp_act = nn.GELU(approximate="tanh")
|
||||
self.modulation = Modulation(hidden_size, double=False)
|
||||
|
||||
def forward(self, x: Tensor, vec: Tensor, pe: Tensor, mask: Tensor = None) -> Tensor:
|
||||
mod, _ = self.modulation(vec)
|
||||
x_mod = (1 + mod.scale) * self.pre_norm(x) + mod.shift
|
||||
qkv, mlp = torch.split(self.linear1(x_mod), [3 * self.hidden_size, self.mlp_hidden_dim], dim=-1)
|
||||
|
||||
q, k, v = rearrange(qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
|
||||
q, k = self.norm(q, k, v)
|
||||
if mask is not None:
|
||||
mask = repeat(mask, 'B L S-> B H L S', H=self.num_heads)
|
||||
# compute attention
|
||||
attn = attention(q, k, v, pe=pe, mask = mask)
|
||||
# compute activation in mlp stream, cat again and run second linear layer
|
||||
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
|
||||
return x + mod.gate * output
|
||||
|
||||
|
||||
class LastLayer(nn.Module):
|
||||
def __init__(self, hidden_size: int, patch_size: int, out_channels: int):
|
||||
super().__init__()
|
||||
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True)
|
||||
self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True))
|
||||
|
||||
def forward(self, x: Tensor, vec: Tensor) -> Tensor:
|
||||
shift, scale = self.adaLN_modulation(vec).chunk(2, dim=1)
|
||||
x = (1 + scale[:, None, :]) * self.norm_final(x) + shift[:, None, :]
|
||||
x = self.linear(x)
|
||||
return x
|
||||
@@ -60,8 +60,8 @@ class FinalLayer(nn.Module):
|
||||
self.linear = nn.Linear(hidden_size,
|
||||
patch_size * patch_size * out_channels,
|
||||
bias=True)
|
||||
self.adaLN_modulation = nn.Sequential(
|
||||
nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True))
|
||||
self.adaLN_modulation = nn.Sequential(nn.SiLU(),
|
||||
nn.Linear(hidden_size, 2 * hidden_size, bias=True))
|
||||
|
||||
def forward(self, x, c):
|
||||
shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)
|
||||
|
||||
@@ -9,7 +9,6 @@ import math
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
|
||||
from scepter.modules.model.backbone.transformer.attention import drop_path
|
||||
|
||||
|
||||
@@ -152,7 +151,6 @@ class SizeEmbedder(TimestepEmbedder):
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
# 返回模型参数的数据类型
|
||||
return next(self.parameters()).dtype
|
||||
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@ from itertools import repeat as iter_repeat
|
||||
from typing import Iterable
|
||||
|
||||
import numpy as np
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
@@ -117,7 +118,6 @@ def apply_2d_rope(xq,
|
||||
# xq_.shape = [b, seq_len, dim // 2, 2]
|
||||
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 2)
|
||||
xk_ = xk.float().reshape(*xk.shape[:-1], -1, 2)
|
||||
# 转为复数域
|
||||
xq_ = torch.view_as_complex(
|
||||
xq_) # [b, seq_len, dim // 2, 2]=>xq.shape = [b, seq_len, dim]
|
||||
xk_ = torch.view_as_complex(xk_)
|
||||
|
||||
@@ -43,7 +43,13 @@ class BaseModel(nn.Module):
|
||||
self._dist_data[key][k] += v
|
||||
else:
|
||||
self._dist_data[key][k] = v
|
||||
|
||||
def collect_probe(self):
|
||||
probe_data_dict = self._probe_data
|
||||
for k, v in self._modules.items():
|
||||
if isinstance(getattr(self, k), BaseModel):
|
||||
for kk, vv in getattr(self, k).collect_probe().items():
|
||||
probe_data_dict[f'{k}/{kk}'] = vv
|
||||
return probe_data_dict
|
||||
def probe_data(self):
|
||||
gather_probe_data = gather_data(self._probe_data)
|
||||
_dist_data_list = gather_data([self._dist_data])
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
from .samplers import BaseDiffusionSampler, FlowEluerSampler, DDIMSampler
|
||||
from .schedules import BaseNoiseScheduler, ScaledLinearScheduler, FlowMatchShiftScheduler
|
||||
from .diffusions import BaseDiffusion, DiffusionFluxRF
|
||||
@@ -0,0 +1,264 @@
|
||||
import os
|
||||
import math
|
||||
import torch
|
||||
from collections import OrderedDict
|
||||
|
||||
from scepter.modules.utils.config import dict_to_yaml, Config
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.modules.model.registry import DIFFUSIONS, NOISE_SCHEDULERS, DIFFUSION_SAMPLERS
|
||||
from tqdm import trange
|
||||
|
||||
@DIFFUSIONS.register_class()
|
||||
class BaseDiffusion(object):
|
||||
para_dict = {
|
||||
"NOISE_SCHEDULER": {},
|
||||
"SAMPLER_SCHEDULER": {},
|
||||
"MIN_SNR_GAMMA": {
|
||||
"value": None,
|
||||
"description": "The minimum SNR gamma value for the loss function."
|
||||
},
|
||||
"PREDICTION_TYPE": {
|
||||
"value": "eps",
|
||||
"description": "The type of prediction to use for the loss function."
|
||||
}
|
||||
}
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(BaseDiffusion, self).__init__()
|
||||
self.logger = logger
|
||||
self.cfg = cfg
|
||||
self.init_params()
|
||||
|
||||
def init_params(self):
|
||||
self.min_snr_gamma = self.cfg.get("MIN_SNR_GAMMA", None)
|
||||
self.prediction_type = self.cfg.get("PREDICTION_TYPE", "eps")
|
||||
self.noise_scheduler = NOISE_SCHEDULERS.build(self.cfg.NOISE_SCHEDULER, logger=self.logger)
|
||||
self.sampler_scheduler = NOISE_SCHEDULERS.build(self.cfg.get("SAMPLER_SCHEDULER", self.cfg.NOISE_SCHEDULER),
|
||||
logger=self.logger)
|
||||
self.num_timesteps = self.noise_scheduler.num_timesteps
|
||||
if self.cfg.have("WORK_DIR") and we.rank == 0:
|
||||
schedule_visualization = os.path.join(self.cfg.WORK_DIR, "noise_schedule.png")
|
||||
with FS.put_to(schedule_visualization) as local_path:
|
||||
self.noise_scheduler.plot_noise_sampling_map(local_path)
|
||||
schedule_visualization = os.path.join(self.cfg.WORK_DIR, "sampler_schedule.png")
|
||||
with FS.put_to(schedule_visualization) as local_path:
|
||||
self.sampler_scheduler.plot_noise_sampling_map(local_path)
|
||||
|
||||
|
||||
def sample(self, noise, model, model_kwargs={}, steps=20, sampler=None, use_dynamic_cfg=False, guide_scale=None, guide_rescale=None,
|
||||
show_progress=False, return_intermediate=None, intermediate_callback=None):
|
||||
assert isinstance(steps, (int, torch.LongTensor))
|
||||
assert return_intermediate in (None, 'x0', 'xt')
|
||||
assert isinstance(sampler, (str, dict, Config))
|
||||
intermediates = []
|
||||
|
||||
def callback_fn(x_t, t, sigma=None, alpha=None):
|
||||
t = t.repeat(len(x_t)).round().long().to(x_t.device)
|
||||
if guide_scale is None or guide_scale == 1.0:
|
||||
out = model(x=x_t, t=t, **model_kwargs)
|
||||
else:
|
||||
if use_dynamic_cfg:
|
||||
guidance_scale = 1 + guide_scale * ((1 - math.cos(math.pi * ((steps - t.item()) / steps) ** 5.0)) / 2)
|
||||
else:
|
||||
guidance_scale = guide_scale
|
||||
y_out = model(x=x_t, t=t, **model_kwargs[0])
|
||||
u_out = model(x=x_t, t=t, **model_kwargs[1])
|
||||
out = u_out + guidance_scale * (y_out - u_out)
|
||||
if guide_rescale is not None and guide_rescale > 0.0:
|
||||
ratio = (
|
||||
y_out.flatten(1).std(dim=1) /
|
||||
(out.flatten(1).std(dim=1) + 1e-12)).view((-1, ) + (1, ) *
|
||||
(y_out.ndim - 1))
|
||||
out *= guide_rescale * ratio + (1 - guide_rescale) * 1.0
|
||||
|
||||
if self.prediction_type == 'x0':
|
||||
x0 = out
|
||||
elif self.prediction_type == 'eps':
|
||||
x0 = (x_t - sigma * out) / alpha
|
||||
elif self.prediction_type == 'v':
|
||||
x0 = alpha * x_t - sigma * out
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f'prediction_type {self.prediction_type} not implemented')
|
||||
|
||||
# print("torch.sum(y_out):", torch.sum(y_out), "torch.sum(u_out):", torch.sum(u_out), "torch.sum(out):",
|
||||
# torch.sum(out), "torch.sum(x0):", torch.sum(x0), "sigmas", sigma, "alphas", alpha)
|
||||
return x0
|
||||
|
||||
sampler_ins = self.get_sampler(sampler)
|
||||
|
||||
# this is ignored for schnell
|
||||
sampler_output = sampler_ins.preprare_sampler(
|
||||
noise,
|
||||
steps=steps,
|
||||
prediction_type=self.prediction_type,
|
||||
scheduler_ins=self.sampler_scheduler,
|
||||
callback_fn=callback_fn
|
||||
)
|
||||
|
||||
for _ in trange(steps, disable=not show_progress):
|
||||
trange.desc = sampler_output.msg
|
||||
sampler_output = sampler_ins.step(sampler_output)
|
||||
if return_intermediate == 'x_0':
|
||||
intermediates.append(sampler_output.x_0)
|
||||
elif return_intermediate == 'x_t':
|
||||
intermediates.append(sampler_output.x_t)
|
||||
if intermediate_callback is not None:
|
||||
intermediate_callback(intermediates[-1])
|
||||
return (sampler_output.x_0, intermediates) if return_intermediate is not None else sampler_output.x_0
|
||||
|
||||
|
||||
def loss(self, x_0, model, model_kwargs={}, reduction='mean', noise=None, **kwargs):
|
||||
# use noise scheduler to add noise
|
||||
if noise is None:
|
||||
noise = torch.randn_like(x_0)
|
||||
schedule_output = self.noise_scheduler.add_noise(x_0, noise)
|
||||
x_t, t, sigma, alpha = schedule_output.x_t, schedule_output.t, schedule_output.sigma, schedule_output.alpha
|
||||
out = model(x=x_t, t=t, **model_kwargs)
|
||||
|
||||
# mse loss
|
||||
target = {
|
||||
'eps': noise,
|
||||
'x0': x_0,
|
||||
'v': alpha * noise - sigma * x_0
|
||||
}[self.prediction_type]
|
||||
|
||||
loss = (out - target).pow(2)
|
||||
if reduction == 'mean':
|
||||
loss = loss.flatten(1).mean(dim=1)
|
||||
|
||||
if self.min_snr_gamma is not None:
|
||||
alphas = self.noise_scheduler.alphas.to(x_0.device)[t]
|
||||
sigmas = self.noise_scheduler.sigmas.pow(2).to(x_0.device)[t]
|
||||
snrs = (alphas / sigmas).clamp(min=1e-20)
|
||||
min_snrs = snrs.clamp(max=self.min_snr_gamma)
|
||||
weights = min_snrs / snrs
|
||||
else:
|
||||
weights = 1
|
||||
|
||||
loss = loss * weights
|
||||
return loss
|
||||
|
||||
def get_sampler(self, sampler):
|
||||
if isinstance(sampler, str):
|
||||
if sampler not in DIFFUSION_SAMPLERS.class_map:
|
||||
if self.logger is not None:
|
||||
self.logger.info(f"{sampler} not in the defined samplers list {DIFFUSION_SAMPLERS.class_map.keys()}")
|
||||
else:
|
||||
print(f"{sampler} not in the defined samplers list {DIFFUSION_SAMPLERS.class_map.keys()}")
|
||||
return None
|
||||
sampler_cfg = Config(cfg_dict={"NAME": sampler}, load=False)
|
||||
sampler_ins = DIFFUSION_SAMPLERS.build(sampler_cfg, logger=self.logger)
|
||||
elif isinstance(sampler, (Config, dict, OrderedDict)):
|
||||
if isinstance(sampler, (dict, OrderedDict)):
|
||||
sampler = Config(cfg_dict={k.upper():v for k, v in dict(sampler).items()}, load=False)
|
||||
sampler_ins = DIFFUSION_SAMPLERS.build(sampler, logger=self.logger)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
return sampler_ins
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f'{self.__class__.__name__}' + ' ' + super().__repr__()
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('DIFFUSIONS',
|
||||
__class__.__name__,
|
||||
BaseDiffusion.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@DIFFUSIONS.register_class()
|
||||
class DiffusionFluxRF(BaseDiffusion):
|
||||
para_dict = {
|
||||
"PREDICTION_TYPE": {
|
||||
"value": "raw",
|
||||
"description": "The type of prediction to use for the loss function."
|
||||
}
|
||||
}
|
||||
para_dict.update(BaseDiffusion.para_dict)
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(DiffusionFluxRF, self).__init__(cfg, logger=logger)
|
||||
self.prediction_type = self.cfg.get("PREDICTION_TYPE", "raw")
|
||||
|
||||
def loss(self, x_0, model, model_kwargs={}, reduction='mean', noise=None, **kwargs):
|
||||
if noise is None:
|
||||
noise = torch.randn_like(x_0)
|
||||
schedule_output = self.noise_scheduler.add_noise(x_0, noise)
|
||||
x_t, t, sigma = schedule_output.x_t, schedule_output.t, schedule_output.sigma
|
||||
out = model(x=x_t, t=sigma, **model_kwargs)
|
||||
# raw
|
||||
if self.prediction_type == "raw":
|
||||
target = noise - x_0
|
||||
out = out
|
||||
elif self.prediction_type == "sigma_scaled":
|
||||
target = x_0
|
||||
out = out * (-sigma) + x_t
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
loss = (target - out) ** 2
|
||||
if reduction == 'mean':
|
||||
loss = loss.flatten(1).mean(dim=1)
|
||||
|
||||
if self.min_snr_gamma is not None:
|
||||
alphas = self.noise_scheduler.alphas.to(x_0.device)[t]
|
||||
sigmas = self.noise_scheduler.sigmas.pow(2).to(x_0.device)[t]
|
||||
snrs = (alphas / sigmas).clamp(min=1e-20)
|
||||
min_snrs = snrs.clamp(max=self.min_snr_gamma)
|
||||
weights = min_snrs / snrs
|
||||
else:
|
||||
weights = 1
|
||||
|
||||
loss = loss * weights
|
||||
return loss
|
||||
|
||||
@torch.no_grad()
|
||||
def sample(self,
|
||||
noise,
|
||||
model,
|
||||
model_kwargs={},
|
||||
steps=20,
|
||||
sampler = None,
|
||||
show_progress=False,
|
||||
return_intermediate=None,
|
||||
intermediate_callback=None,
|
||||
**kwargs):
|
||||
# sanity check
|
||||
assert isinstance(steps, (int, torch.LongTensor))
|
||||
assert return_intermediate in (None, 'x0', 'xt')
|
||||
assert isinstance(sampler, (str, dict, Config))
|
||||
intermediates = []
|
||||
|
||||
def callback_fn(x_t, t, sigma):
|
||||
sigma = torch.full((x_t.shape[0],), sigma, dtype=x_t.dtype, device=x_t.device)
|
||||
x_0 = model(x = x_t, t=sigma, **model_kwargs)
|
||||
return x_0
|
||||
|
||||
sampler_ins = self.get_sampler(sampler)
|
||||
|
||||
# this is ignored for schnell
|
||||
sampler_output = sampler_ins.preprare_sampler(
|
||||
noise,
|
||||
steps = steps,
|
||||
prediction_type=self.prediction_type,
|
||||
scheduler_ins = self.sampler_scheduler,
|
||||
callback_fn=callback_fn
|
||||
)
|
||||
|
||||
for _ in trange(steps, disable=not show_progress):
|
||||
trange.desc = sampler_output.msg
|
||||
sampler_output = sampler_ins.step(sampler_output)
|
||||
if return_intermediate == 'x_0':
|
||||
intermediates.append(sampler_output.x_0)
|
||||
elif return_intermediate == 'x_t':
|
||||
intermediates.append(sampler_output.x_t)
|
||||
if intermediate_callback is not None:
|
||||
intermediate_callback(intermediates[-1])
|
||||
return (sampler_output.x_0, intermediates) if return_intermediate is not None else sampler_output.x_t
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('DIFFUSIONS',
|
||||
__class__.__name__,
|
||||
DiffusionFluxRF.para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,211 @@
|
||||
from dataclasses import dataclass, field
|
||||
import torch
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.model.registry import DIFFUSION_SAMPLERS
|
||||
from .util import _i
|
||||
|
||||
@dataclass
|
||||
class SamplerOutput(object):
|
||||
callback_fn: callable
|
||||
prediction_type: str
|
||||
alphas: torch.Tensor
|
||||
betas: torch.Tensor
|
||||
sigmas: torch.Tensor
|
||||
alphas_init: torch.Tensor
|
||||
betas_init: torch.Tensor
|
||||
sigmas_init: torch.Tensor
|
||||
ts: torch.Tensor
|
||||
x_t: torch.Tensor
|
||||
x_0: torch.Tensor
|
||||
step: int
|
||||
msg: str
|
||||
|
||||
def add_custom_field(self, key: str, value) -> None:
|
||||
self.__setattr__(key, value)
|
||||
|
||||
|
||||
@DIFFUSION_SAMPLERS.register_class("base")
|
||||
class BaseDiffusionSampler(object):
|
||||
para_dict = {}
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(BaseDiffusionSampler, self).__init__()
|
||||
self.logger = logger
|
||||
self.cfg = cfg
|
||||
self.init_params()
|
||||
|
||||
def init_params(self):
|
||||
self.discretization_type = self.cfg.get("DISCRETIZATION_TYPE", "linspace")
|
||||
self.discard_penultimate_step = self.cfg.get("DISCARD_PENULTIMATE_STEP", False)
|
||||
self.free_steps = self.cfg.get("FREE_STEPS", None)
|
||||
self.t_max = self.cfg.get("T_MAX", None)
|
||||
self.t_min = self.cfg.get("T_MIN", None)
|
||||
|
||||
def discretization(self, steps=20, num_timesteps=1000, **kwargs):
|
||||
# get timesteps
|
||||
if isinstance(steps, int):
|
||||
steps += 1 if self.discard_penultimate_step else 0
|
||||
t_max = num_timesteps - 1 if self.t_max is None else self.t_max
|
||||
t_min = 0 if self.t_min is None else self.t_min
|
||||
|
||||
# discretize timesteps
|
||||
if self.discretization_type == 'leading':
|
||||
steps = torch.arange(t_min, t_max + 1,
|
||||
(t_max - t_min + 1) / steps).flip(0)
|
||||
elif self.discretization_type == 'linspace':
|
||||
steps = torch.linspace(t_max, t_min, steps)
|
||||
elif self.discretization_type == 'trailing':
|
||||
steps = torch.arange(t_max, t_min - 1,
|
||||
-((t_max - t_min + 1) / steps))
|
||||
elif self.discretization_type == 'free':
|
||||
steps = torch.tensor(self.free_steps)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f'{self.discretization_type} discretization not implemented')
|
||||
steps = steps.clamp_(t_min, t_max)
|
||||
elif isinstance(steps, list):
|
||||
steps = torch.tensor(steps)
|
||||
timesteps = torch.as_tensor(steps, dtype=torch.float32)
|
||||
return timesteps
|
||||
|
||||
def preprare_sampler(self, noise, steps=20, scheduler_ins=None, prediction_type="",
|
||||
sigmas=None, betas=None, alphas=None, callback_fn = None,
|
||||
**kwargs):
|
||||
'''
|
||||
1. Control the model's inputs and outputs externally in the solver by callback_fn,
|
||||
and perform the conversion between x0 and xt internally within the solver.
|
||||
2. The function callback_fn use the x0 and xt as the default inputs and also give me
|
||||
the x0 and xt as output. The other inputs will be set in kwargs.
|
||||
3. The basic parameters of the schedule should be set manually.
|
||||
4. To ensure the safety of threading, use the instance of SamplerOutput as the manager,
|
||||
which manage all necessary information.
|
||||
'''
|
||||
num_timesteps = scheduler_ins.num_timesteps if scheduler_ins is not None else 1000
|
||||
timestamps = self.discretization(steps, num_timesteps=num_timesteps, **kwargs)
|
||||
alphas = scheduler_ins.t_to_alpha(timestamps, **kwargs) if scheduler_ins is not None else alphas
|
||||
betas = scheduler_ins.t_to_beta(timestamps, **kwargs) if scheduler_ins is not None else betas
|
||||
sigmas = scheduler_ins.t_to_sigma(timestamps, **kwargs) if scheduler_ins is not None else sigmas
|
||||
alphas_init = scheduler_ins.t_to_alpha_init(timestamps, **kwargs) if scheduler_ins is not None else alphas
|
||||
betas_init = scheduler_ins.t_to_beta_init(timestamps, **kwargs) if scheduler_ins is not None else betas
|
||||
sigmas_init = scheduler_ins.t_to_sigma_init(timestamps, **kwargs) if scheduler_ins is not None else sigmas
|
||||
|
||||
output = SamplerOutput(
|
||||
callback_fn=callback_fn,
|
||||
prediction_type=prediction_type,
|
||||
alphas=alphas,
|
||||
betas=betas,
|
||||
sigmas=sigmas,
|
||||
alphas_init=alphas_init,
|
||||
betas_init=betas_init,
|
||||
sigmas_init=sigmas_init,
|
||||
ts=timestamps,
|
||||
x_t=noise,
|
||||
x_0=noise,
|
||||
step=0,
|
||||
msg=f"step 0"
|
||||
)
|
||||
return output
|
||||
|
||||
def step(self, sampler_ouput):
|
||||
raise NotImplementedError(f'DiffusionSampler step function not implemented')
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f'{self.__class__.__name__}' + ' ' + super().__repr__()
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('DIFFUSION_SAMPLERS',
|
||||
__class__.__name__,
|
||||
BaseDiffusionSampler.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@DIFFUSION_SAMPLERS.register_class("eluer")
|
||||
class EulerSampler(BaseDiffusionSampler):
|
||||
def step(self, sampler_ouput):
|
||||
pass
|
||||
|
||||
|
||||
@DIFFUSION_SAMPLERS.register_class("ddim")
|
||||
class DDIMSampler(BaseDiffusionSampler):
|
||||
|
||||
def init_params(self):
|
||||
super().init_params()
|
||||
self.eta = self.cfg.get('ETA', 0.)
|
||||
self.discretization_type = self.cfg.get("DISCRETIZATION_TYPE", "trailing")
|
||||
|
||||
def preprare_sampler(self, noise, steps=20, scheduler_ins=None, prediction_type="",
|
||||
sigmas=None, betas=None, alphas=None, callback_fn = None,
|
||||
**kwargs):
|
||||
output = super().preprare_sampler(noise, steps, scheduler_ins, prediction_type, sigmas, betas, alphas, callback_fn, **kwargs)
|
||||
sigmas = output.sigmas
|
||||
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
|
||||
sigmas_vp = (sigmas**2 / (1 + sigmas**2))**0.5
|
||||
sigmas_vp[sigmas == float('inf')] = 1.
|
||||
output.add_custom_field('sigmas_vp', sigmas_vp)
|
||||
return output
|
||||
|
||||
def step(self, sampler_output):
|
||||
step = sampler_output.step
|
||||
x_t = sampler_output.x_t
|
||||
t = sampler_output.ts[step]
|
||||
sigmas_vp = sampler_output.sigmas_vp.to(x_t.device)
|
||||
alpha_init = _i(sampler_output.alphas_init, step, x_t)
|
||||
sigma_init = _i(sampler_output.sigmas_init, step, x_t)
|
||||
|
||||
x = sampler_output.callback_fn(x_t, t, sigma_init, alpha_init)
|
||||
noise_factor = self.eta * (sigmas_vp[step + 1] ** 2 / sigmas_vp[step] ** 2 *
|
||||
(1 - (1 - sigmas_vp[step] ** 2) /
|
||||
(1 - sigmas_vp[step + 1] ** 2)))
|
||||
d = (x_t - (1 - sigmas_vp[step] ** 2) ** 0.5 * x) / sigmas_vp[step]
|
||||
x = (1 - sigmas_vp[step + 1] ** 2) ** 0.5 * x + \
|
||||
(sigmas_vp[step + 1] ** 2 - noise_factor ** 2) ** 0.5 * d
|
||||
sampler_output.x_0 = x
|
||||
if sigmas_vp[step + 1] > 0:
|
||||
x += noise_factor * torch.randn_like(x)
|
||||
# print("i:", step, "sigma_init:", sigma_init, "alpha_init", alpha_init, "sigmas_vp[i]", sigmas_vp[step], "torch.sum(x_0):", torch.sum(x_0), "torch.sum(x):", torch.sum(x))
|
||||
sampler_output.x_t = x
|
||||
sampler_output.step += 1
|
||||
sampler_output.msg = f'step {step}'
|
||||
return sampler_output
|
||||
|
||||
|
||||
@DIFFUSION_SAMPLERS.register_class("flow_eluer")
|
||||
class FlowEluerSampler(BaseDiffusionSampler):
|
||||
def preprare_sampler(self, noise, steps=20, scheduler_ins=None, prediction_type="",
|
||||
sigmas=None, betas=None, alphas=None, callback_fn = None,
|
||||
**kwargs):
|
||||
if noise.ndim == 3:
|
||||
seq_len = noise.shape[2] // 4
|
||||
else:
|
||||
n, _, h, w = noise.shape
|
||||
seq_len = (h // 2 * w // 2)
|
||||
kwargs["seq_len"] = seq_len
|
||||
output = super().preprare_sampler(noise, steps, scheduler_ins, prediction_type, sigmas, betas, alphas, callback_fn, **kwargs)
|
||||
return output
|
||||
|
||||
def step(self, sampler_output):
|
||||
step = sampler_output.step
|
||||
x_t = sampler_output.x_t
|
||||
sigma_curr, sigma_prev = sampler_output.sigmas[step], sampler_output.sigmas[step + 1]
|
||||
prediction_type = sampler_output.prediction_type
|
||||
assert prediction_type in ("raw", "sigma_scaled")
|
||||
t = sampler_output.ts[step]
|
||||
x_0 = sampler_output.callback_fn(x_t, t, sigma_curr)
|
||||
x_t = x_t + (sigma_prev - sigma_curr) * x_0
|
||||
sampler_output.x_0 = x_0
|
||||
sampler_output.x_t = x_t
|
||||
sampler_output.step += 1
|
||||
sampler_output.msg = f'step {step}, sigma_curr: {sigma_curr}, sigma_prev: {sigma_prev}'
|
||||
return sampler_output
|
||||
|
||||
def discretization(self, steps=20, num_timesteps = 1000, **kwargs):
|
||||
# extra step for zero
|
||||
timesteps = torch.linspace(num_timesteps, 0, steps + 1)
|
||||
return timesteps
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('DIFFUSION_SAMPLERS',
|
||||
__class__.__name__,
|
||||
FlowEluerSampler.para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,533 @@
|
||||
import math
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.math_plot import plot_multi_curves
|
||||
from scepter.modules.model.registry import NOISE_SCHEDULERS
|
||||
from torch import Tensor
|
||||
|
||||
from .util import _i
|
||||
|
||||
@dataclass
|
||||
class ScheduleOutput(object):
|
||||
x_t: torch.Tensor
|
||||
x_0: torch.Tensor
|
||||
t: torch.Tensor
|
||||
sigma: torch.Tensor
|
||||
alpha: torch.Tensor
|
||||
custom_fields: dict = field(default_factory=dict)
|
||||
|
||||
def add_custom_field(self, key: str, value) -> None:
|
||||
self.__setattr__(key, value)
|
||||
|
||||
|
||||
@NOISE_SCHEDULERS.register_class()
|
||||
class BaseNoiseScheduler(object):
|
||||
'''
|
||||
In the diffusion model, the parameters related to the noise schedule are alpha, beta,
|
||||
and sigma. The following are the definitions of the above three parameters, which should
|
||||
be the basic property for the instance of noise scheduler.
|
||||
\alpha_{t} = \sqrt{1 - \beta_{t}^2} \alpha is the strength of signal and \beta is the strength of noise
|
||||
\sigma_{t} = \sqrt{1 - \overline\alpha} = \sqrt{1 - \prod_{i=1}^{t}\alpha^2_{i}} (P(x_{t}|x_{0}) ~ N(\overline\alpha x_{0}, \sigma^2))
|
||||
|
||||
where sigma_{t} is the var of p(x_{t-1}|x_{t}, x_{0}).
|
||||
|
||||
(reference to https://arxiv.org/abs/2010.02502)
|
||||
let sigma transfer to beta:
|
||||
square_\beta = 1 - \frac{1 - square_\sigma_{t}}{1 - square_\sigma_{t - 1 }}
|
||||
|
||||
'''
|
||||
para_dict = {
|
||||
"NUM_TIMESTEPS": {
|
||||
"value": 1000,
|
||||
"description": "The number of timesteps for sampling."
|
||||
},
|
||||
}
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(BaseNoiseScheduler, self).__init__()
|
||||
self.logger = logger
|
||||
self.cfg = cfg
|
||||
self.init_params()
|
||||
self.get_schedule()
|
||||
# self.check_function()
|
||||
|
||||
def init_params(self):
|
||||
self.num_timesteps = self.cfg.get("NUM_TIMESTEPS", 1000)
|
||||
self._sample_steps = torch.arange(self.num_timesteps, dtype=torch.float32)
|
||||
self._sigmas, self._betas, self._alphas, self._timesteps = None, None, None, None
|
||||
def check_function(self):
|
||||
# for the same t, we should gurantee t_to_sigma and sigma_to_t is aligned
|
||||
try:
|
||||
predict_timestamps = self.sigma_to_t(self.sigmas)
|
||||
predict_sigmas = self.t_to_sigma(self._timesteps)
|
||||
diff_sigmas = torch.sum(torch.abs(predict_sigmas - self.sigmas))
|
||||
diff_timestamps = torch.sum(torch.abs(predict_timestamps - self._timesteps))
|
||||
if diff_sigmas > 1e-3 or diff_timestamps > 1:
|
||||
self.logger.info(f"The noise scheduler {self.__class__.__name__} is not correct, "
|
||||
f"please check the function sigma_to_t or t_to_sigma."
|
||||
f"Info: diff sigmas {diff_sigmas}, diff timestamps {diff_timestamps}")
|
||||
raise "The noise scheduler checked failed."
|
||||
else:
|
||||
self.logger.info(f"The noise scheduler {self.__class__.__name__} is checked and passed.")
|
||||
except Exception as e:
|
||||
if isinstance(e, NotImplementedError):
|
||||
self.logger.info("Not implemented function sigma_to_t or t_to_sigma, skip check.")
|
||||
else:
|
||||
self.logger.info(f"The noise scheduler {self.__class__.__name__} is not correct, "
|
||||
f"please check the function sigma_to_t or t_to_sigma. Error: {e}")
|
||||
raise e
|
||||
|
||||
def get_schedule(self):
|
||||
raise NotImplementedError(f'NoiseScheduler get_schedule function not implemented')
|
||||
|
||||
def square_betas_to_sigmas(self, square_betas):
|
||||
return torch.sqrt(1 - torch.cumprod(1 - square_betas, dim=0))
|
||||
def sigmas_to_square_betas(self, sigmas):
|
||||
square_alphas = 1 - sigmas ** 2
|
||||
betas = 1 - torch.cat([square_alphas[:1], square_alphas[1:] / square_alphas[:-1]])
|
||||
return betas
|
||||
def sigma_to_t(self, sigma, **kwargs):
|
||||
if sigma == float('inf'):
|
||||
t = torch.full_like(sigma, len(self._sigmas) - 1)
|
||||
else:
|
||||
log_sigmas = torch.sqrt(self._sigmas**2 /
|
||||
(1 - self._sigmas**2)).log().to(sigma)
|
||||
log_sigma = sigma.log()
|
||||
dists = log_sigma - log_sigmas[:, None]
|
||||
low_idx = dists.ge(0).cumsum(dim=0).argmax(dim=0).clamp(
|
||||
max=log_sigmas.shape[0] - 2)
|
||||
high_idx = low_idx + 1
|
||||
low, high = log_sigmas[low_idx], log_sigmas[high_idx]
|
||||
w = (low - log_sigma) / (low - high)
|
||||
w = w.clamp(0, 1)
|
||||
t = (1 - w) * low_idx + w * high_idx
|
||||
t = t.view(sigma.shape)
|
||||
if t.ndim == 0:
|
||||
t = t.unsqueeze(0)
|
||||
return t
|
||||
def t_to_sigma(self, t, **kwargs):
|
||||
t = t.float()
|
||||
low_idx, high_idx, w = t.floor().long(), t.ceil().long(), t.frac()
|
||||
log_sigmas = torch.sqrt(self.sigmas**2 /
|
||||
(1 - self.sigmas**2)).log().to(t)
|
||||
log_sigma = (1 - w) * log_sigmas[low_idx] + w * log_sigmas[high_idx]
|
||||
log_sigma[torch.isnan(log_sigma)
|
||||
| torch.isinf(log_sigma)] = float('inf')
|
||||
return log_sigma.exp()
|
||||
|
||||
def t_to_alpha(self, t, **kwargs):
|
||||
sigma = self.t_to_sigma(t)
|
||||
square_beta = self.sigmas_to_square_betas(sigma)
|
||||
return torch.sqrt(1 - square_beta)
|
||||
|
||||
def t_to_beta(self, t, **kwargs):
|
||||
sigma = self.t_to_sigma(t)
|
||||
square_beta = self.sigmas_to_square_betas(sigma)
|
||||
return torch.sqrt(square_beta)
|
||||
|
||||
def add_noise(self, x_0, noise = None, t = None):
|
||||
if t is None:
|
||||
t = torch.randint(0, self.num_timesteps, (x_0.shape[0],), device=x_0.device).long()
|
||||
alpha = _i(self.alphas, t, x_0)
|
||||
sigma = _i(self.sigmas, t, x_0)
|
||||
x_t = alpha * x_0 + sigma * noise
|
||||
|
||||
return ScheduleOutput(x_0 = x_0, x_t = x_t, t = t, alpha=alpha, sigma=sigma)
|
||||
|
||||
def t_to_alpha_init(self, t, **kwargs):
|
||||
indices = t.long()
|
||||
indices[indices >= self.num_timesteps] = self.num_timesteps - 1
|
||||
timesteps = self.timesteps.to(t)[indices]
|
||||
step_indices = [(self.timesteps.to(t) == t).nonzero().item() for t in timesteps]
|
||||
alpha = self.alphas[step_indices].flatten().to(t)
|
||||
return alpha
|
||||
|
||||
def t_to_beta_init(self, t, **kwargs):
|
||||
indices = t.long()
|
||||
indices[indices >= self.num_timesteps] = self.num_timesteps - 1
|
||||
timesteps = self.timesteps.to(t)[indices]
|
||||
step_indices = [(self.timesteps.to(t) == t).nonzero().item() for t in timesteps]
|
||||
beta = self.betas[step_indices].flatten().to(t)
|
||||
return beta
|
||||
|
||||
def t_to_sigma_init(self, t, **kwargs):
|
||||
indices = t.long()
|
||||
indices[indices >= self.num_timesteps] = self.num_timesteps - 1
|
||||
timesteps = self.timesteps.to(t)[indices]
|
||||
step_indices = [(self.timesteps.to(t) == t).nonzero().item() for t in timesteps]
|
||||
sigma = self.sigmas[step_indices].flatten().to(t)
|
||||
return sigma
|
||||
|
||||
def rescale_zero_terminal_snr(self, alphas_cumprod):
|
||||
"""
|
||||
Rescales betas to have zero terminal SNR Based on https://arxiv.org/pdf/2305.08891.pdf (Algorithm 1)
|
||||
Args:
|
||||
betas (`torch.Tensor`):
|
||||
the betas that the scheduler is being initialized with.
|
||||
Returns:
|
||||
`torch.Tensor`: rescaled betas with zero terminal SNR
|
||||
"""
|
||||
alphas_bar_sqrt = alphas_cumprod.sqrt()
|
||||
# Store old values.
|
||||
alphas_bar_sqrt_0 = alphas_bar_sqrt[0].clone()
|
||||
alphas_bar_sqrt_T = alphas_bar_sqrt[-1].clone()
|
||||
# Shift so the last timestep is zero.
|
||||
alphas_bar_sqrt -= alphas_bar_sqrt_T
|
||||
# Scale so the first timestep is back to the old value.
|
||||
alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 - alphas_bar_sqrt_T)
|
||||
# Convert alphas_bar_sqrt to betas
|
||||
alphas_bar = alphas_bar_sqrt ** 2 # Revert sqrt
|
||||
return alphas_bar
|
||||
|
||||
@property
|
||||
def sigmas(self):
|
||||
return self._sigmas
|
||||
|
||||
@property
|
||||
def betas(self):
|
||||
return self._betas
|
||||
|
||||
@property
|
||||
def alphas(self):
|
||||
return self._alphas
|
||||
|
||||
@property
|
||||
def timesteps(self):
|
||||
return self._timesteps
|
||||
|
||||
# plot the noise sampling map
|
||||
def plot_noise_sampling_map(self, save_path):
|
||||
y = [
|
||||
{"data": self._sigmas.cpu().numpy(), "label": "sigmas"},
|
||||
{"data": self._betas.cpu().numpy(), "label": "betas"},
|
||||
{"data": self._alphas.cpu().numpy(), "label": "alphas"},
|
||||
{"data": self._timesteps.cpu().numpy()/self.num_timesteps, "label": "timesteps"}
|
||||
]
|
||||
plot_multi_curves(
|
||||
x=self._sample_steps.cpu().numpy(),
|
||||
y=y,
|
||||
x_label='timesteps',
|
||||
y_label=None,
|
||||
title=f"{self.__class__.__name__}'s noise sampling map",
|
||||
save_path=save_path
|
||||
)
|
||||
return save_path
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f'{self.__class__.__name__}' + ' ' + super().__repr__()
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('NOISE_SCHEDULER',
|
||||
__class__.__name__,
|
||||
BaseNoiseScheduler.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@NOISE_SCHEDULERS.register_class()
|
||||
class ScaledLinearScheduler(BaseNoiseScheduler):
|
||||
para_dict = {}
|
||||
def init_params(self):
|
||||
super().init_params()
|
||||
self.beta_min = self.cfg.get('BETA_MIN', 0.00085)
|
||||
self.beta_max = self.cfg.get('BETA_MAX', 0.012)
|
||||
self.snr_shift_scale = self.cfg.get('SNR_SHIFT_SCALE', None)
|
||||
self.rescale_betas_zero_snr = self.cfg.get('RESCALE_BETAS_ZERO_SNR', False)
|
||||
|
||||
def square_betas_to_sigmas(self, square_betas, snr_shift_scale=None, rescale_betas_zero_snr=False):
|
||||
if snr_shift_scale is not None or rescale_betas_zero_snr:
|
||||
alphas_cumprod = torch.cumprod(1 - square_betas, dim=0)
|
||||
if snr_shift_scale is not None and snr_shift_scale > 0:
|
||||
alphas_cumprod = alphas_cumprod / (snr_shift_scale + (1 - snr_shift_scale) * alphas_cumprod)
|
||||
if rescale_betas_zero_snr:
|
||||
alphas_cumprod = self.rescale_zero_terminal_snr(alphas_cumprod)
|
||||
return torch.sqrt(1 - alphas_cumprod)
|
||||
else:
|
||||
return torch.sqrt(1 - torch.cumprod(1 - square_betas, dim=0))
|
||||
|
||||
def get_schedule(self):
|
||||
square_betas = torch.linspace(self.beta_min**0.5, self.beta_max**0.5, self.num_timesteps, dtype=torch.float32) ** 2
|
||||
self._sigmas = self.square_betas_to_sigmas(square_betas, self.snr_shift_scale, self.rescale_betas_zero_snr)
|
||||
self._betas = torch.sqrt(square_betas)
|
||||
self._alphas = torch.sqrt(1 - self._sigmas ** 2)
|
||||
self._timesteps = torch.arange(len(self._sigmas), dtype=torch.float32)
|
||||
|
||||
|
||||
@NOISE_SCHEDULERS.register_class()
|
||||
class FlowMatchUniformScheduler(BaseNoiseScheduler):
|
||||
def get_schedule(self):
|
||||
timesteps = np.linspace(1, self.num_timesteps, self.num_timesteps, dtype=np.float32).copy()
|
||||
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
|
||||
self._timesteps = timesteps
|
||||
self._sigmas = self.t_to_sigma(timesteps)
|
||||
self._betas = torch.sqrt(self.sigmas_to_square_betas(self._sigmas))
|
||||
self._alphas = torch.sqrt(1 - self.betas ** 2)
|
||||
|
||||
def add_noise(self, x_0, noise = None, t = None):
|
||||
if t is None:
|
||||
t = torch.rand((x_0.shape[0],), device=x_0.device)
|
||||
sigma = self.t_to_sigma(t)
|
||||
shape = (x_0.size(0),) + (1,) * (x_0.ndim - 1)
|
||||
x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise
|
||||
return ScheduleOutput(x_0 = x_0, x_t = x_t, t = t, sigma=sigma, alpha=self.t_to_alpha(t))
|
||||
|
||||
def sigma_to_t(self, sigma, **kwargs):
|
||||
return sigma * self.num_timesteps
|
||||
|
||||
def t_to_sigma(self, t, **kwargs):
|
||||
return t/self.num_timesteps
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('NOISE_SCHEDULER',
|
||||
__class__.__name__,
|
||||
FlowMatchUniformScheduler.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@NOISE_SCHEDULERS.register_class()
|
||||
class FlowMatchSigmoidScheduler(FlowMatchUniformScheduler):
|
||||
para_dict = {
|
||||
"SIGMOID_SCALE": {
|
||||
"value": 1,
|
||||
"description": "The scale for the sigmoid function."
|
||||
}
|
||||
}
|
||||
def init_params(self):
|
||||
super().init_params()
|
||||
self.sigmoid_scale = self.cfg.get("SIGMOID_SCALE", 1)
|
||||
|
||||
def sigma_to_t(self, sigma, **kwargs):
|
||||
t = - torch.log(1/sigma - 1)/self.sigmoid_scale
|
||||
return t * self.num_timesteps
|
||||
|
||||
def t_to_sigma(self, t, **kwargs):
|
||||
return torch.sigmoid(self.sigmoid_scale * t/self.num_timesteps)
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('NOISE_SCHEDULER',
|
||||
__class__.__name__,
|
||||
FlowMatchSigmoidScheduler.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@NOISE_SCHEDULERS.register_class()
|
||||
class FlowMatchShiftScheduler(FlowMatchUniformScheduler):
|
||||
para_dict = {
|
||||
"SHIFT": {
|
||||
"value": 3,
|
||||
"description": "The shift factor for the timestamp."
|
||||
},
|
||||
"SIGMOID_SCALE": {
|
||||
"value": 1,
|
||||
"description": "The scale for the sigmoid function."
|
||||
}
|
||||
}
|
||||
|
||||
def init_params(self):
|
||||
super().init_params()
|
||||
self.shift = self.cfg.get("SHIFT", 3)
|
||||
self.sigmoid_scale = self.cfg.get("SIGMOID_SCALE", 1)
|
||||
|
||||
def add_noise(self, x_0, noise = None, t = None):
|
||||
if t is None:
|
||||
logits_norm = torch.randn(x_0.shape[0], device=x_0.device)
|
||||
logits_norm = logits_norm * self.sigmoid_scale # larger scale for more uniform sampling
|
||||
t = logits_norm.sigmoid() * self.num_timesteps
|
||||
sigma = self.t_to_sigma(t)
|
||||
shape = (x_0.size(0),) + (1,) * (x_0.ndim - 1)
|
||||
x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise
|
||||
return ScheduleOutput(x_0 = x_0, x_t = x_t, t = t, sigma=sigma, alpha=self.t_to_alpha(t))
|
||||
|
||||
def sigma_to_t(self, sigma, **kwargs):
|
||||
t = sigma/(sigma - self.shift * sigma + self.shift)
|
||||
return t * self.num_timesteps
|
||||
|
||||
def t_to_sigma(self, t, **kwargs):
|
||||
t = t / self.num_timesteps
|
||||
return (t * self.shift) / (1 + (self.shift - 1) * t)
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('NOISE_SCHEDULER',
|
||||
__class__.__name__,
|
||||
FlowMatchShiftScheduler.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@NOISE_SCHEDULERS.register_class()
|
||||
class FlowMatchFluxShiftScheduler(FlowMatchUniformScheduler):
|
||||
para_dict = {
|
||||
"SHIFT": {
|
||||
"value": True,
|
||||
"description": "Use timestamp shift or not, default is True."
|
||||
},
|
||||
"SIGMOID_SCALE": {
|
||||
"value": 1,
|
||||
"description": "The scale of sigmoid function for sampling timesteps."
|
||||
},
|
||||
"BASE_SHIFT": {
|
||||
"value": 0.5,
|
||||
"description": "The base shift factor for the timestamp."
|
||||
},
|
||||
"MAX_SHIFT": {
|
||||
"value": 1.15,
|
||||
"description": "The max shift factor for the timestamp."
|
||||
}
|
||||
}
|
||||
|
||||
def init_params(self):
|
||||
super().init_params()
|
||||
self.shift = self.cfg.get("SHIFT", True)
|
||||
self.sigmoid_scale = self.cfg.get("SIGMOID_SCALE", 1)
|
||||
self.base_shift = self.cfg.get('BASE_SHIFT', 0.5)
|
||||
self.max_shift = self.cfg.get('MAX_SHIFT', 1.15)
|
||||
|
||||
def time_shift(self, mu: float, sigma_scale: float, t: Tensor):
|
||||
return math.exp(mu) / (math.exp(mu) + (1 / t - 1) ** sigma_scale)
|
||||
def sigma_shift(self, mu: float, sigma_scale: float, sigma: Tensor):
|
||||
return 1/(torch.pow((1-sigma) * math.exp(mu)/sigma, sigma_scale) + 1)
|
||||
|
||||
def get_lin_function(self,
|
||||
x1: float = 256, y1: float = 0.5, x2: float = 4096, y2: float = 1.15
|
||||
) -> Callable[[float], float]:
|
||||
m = (y2 - y1) / (x2 - x1)
|
||||
b = y1 - m * x1
|
||||
return lambda x: m * x + b
|
||||
|
||||
def add_noise(self, x_0, noise = None, t = None):
|
||||
if x_0.ndim == 3:
|
||||
seq_len = x_0.shape[2] // 4
|
||||
else:
|
||||
n, _, h, w = x_0.shape
|
||||
seq_len = (h // 2 * w // 2)
|
||||
if t is None:
|
||||
logits_norm = torch.randn(x_0.shape[0], device=x_0.device)
|
||||
logits_norm = logits_norm * self.sigmoid_scale # larger scale for more uniform sampling
|
||||
t = logits_norm.sigmoid() * self.num_timesteps
|
||||
sigma = self.t_to_sigma(t, seq_len=seq_len)
|
||||
shape = (x_0.size(0),) + (1,) * (x_0.ndim - 1)
|
||||
x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise
|
||||
return ScheduleOutput(x_0 = x_0, x_t = x_t, t = t, sigma=sigma, alpha=self.t_to_alpha(t))
|
||||
|
||||
def sigma_to_t(self, sigma, **kwargs):
|
||||
seq_len = kwargs.get('seq_len', 256)
|
||||
if self.shift:
|
||||
mu = self.get_lin_function(y1=self.base_shift, y2=self.max_shift)(seq_len)
|
||||
sigma = self.sigma_shift(mu, 1.0, sigma)
|
||||
t = torch.as_tensor(sigma, dtype=torch.float32)
|
||||
return t * self.num_timesteps
|
||||
|
||||
def t_to_sigma(self, t, **kwargs):
|
||||
seq_len = kwargs.get('seq_len', 256)
|
||||
t = t/self.num_timesteps
|
||||
if self.shift:
|
||||
mu = self.get_lin_function(y1=self.base_shift, y2=self.max_shift)(seq_len)
|
||||
t = self.time_shift(mu, 1.0, t)
|
||||
sigma = torch.as_tensor(t, dtype=torch.float32)
|
||||
return sigma
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('NOISE_SCHEDULER',
|
||||
__class__.__name__,
|
||||
FlowMatchFluxShiftScheduler.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@NOISE_SCHEDULERS.register_class()
|
||||
class FlowMatchSigmaScheduler(FlowMatchUniformScheduler):
|
||||
para_dict = {
|
||||
"WEIGHTING_SCHEME" : {
|
||||
"value": "logit_normal",
|
||||
"description": "The weighting scheme for sampling timesteps, "
|
||||
"choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']."
|
||||
},
|
||||
"SHIFT": {
|
||||
"value": 3.0,
|
||||
"description": "The shift factor for the timestamp."
|
||||
},
|
||||
"LOGIT_MEAN" : {
|
||||
"value": 0.0,
|
||||
"description": "The mean of the logit distribution for sampling timesteps."
|
||||
},
|
||||
"LOGIT_STD" : {
|
||||
"value": 1.0,
|
||||
"description": "The standard deviation of the logit distribution for sampling timesteps."
|
||||
},
|
||||
"MODE_SCALE" : {
|
||||
"value": 1.29,
|
||||
"description": "The scale factor for the mode of the logit distribution for sampling timesteps."
|
||||
}
|
||||
}
|
||||
|
||||
def init_params(self):
|
||||
super().init_params()
|
||||
self.weighting_scheme = self.cfg.get("WEIGHTING_SCHEME", "logit_normal")
|
||||
self.logit_mean = self.cfg.get("LOGIT_MEAN", 0.0)
|
||||
self.logit_std = self.cfg.get("LOGIT_STD", 1.0)
|
||||
self.mode_scale = self.cfg.get("MODE_SCALE", 1.29)
|
||||
self.shift = self.cfg.get("SHIFT", 1.0)
|
||||
|
||||
def get_schedule(self):
|
||||
timesteps = np.linspace(1, self.num_timesteps, self.num_timesteps, dtype=np.float32).copy()
|
||||
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
|
||||
self._timesteps = timesteps
|
||||
timesteps = timesteps / self.num_timesteps
|
||||
self._sigmas = self.shift * timesteps / (1 + (self.shift - 1) * timesteps)
|
||||
self._betas = torch.sqrt(self.sigmas_to_square_betas(self._sigmas))
|
||||
self._alphas = torch.sqrt(1 - self.betas ** 2)
|
||||
|
||||
def add_noise(self, x_0, noise=None, t=None):
|
||||
if t is None:
|
||||
if self.weighting_scheme == "logit_normal":
|
||||
t = torch.normal(mean=self.logit_mean, std=self.logit_std, size=(x_0.shape[0],), device=x_0.device)
|
||||
else:
|
||||
t = torch.rand(x_0.shape[0], device=x_0.device)
|
||||
t = self.compute_density_for_timestep_sampling(t) * self.num_timesteps
|
||||
sigma = self.t_to_sigma(t)
|
||||
shape = (x_0.size(0),) + (1,) * (x_0.ndim - 1)
|
||||
x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise
|
||||
return ScheduleOutput(x_0=x_0, x_t=x_t, t=t, sigma=sigma, alpha=self.t_to_alpha(t))
|
||||
|
||||
def compute_density_for_timestep_sampling(self, t):
|
||||
"""Compute the density for sampling the timesteps when doing SD3 training.
|
||||
Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528.
|
||||
SD3 paper reference: https://arxiv.org/abs/2403.03206v1.
|
||||
"""
|
||||
if self.weighting_scheme == "logit_normal":
|
||||
# See 3.1 in the SD3 paper ($rf/lognorm(0.00,1.00)$).
|
||||
t = torch.nn.functional.sigmoid(t)
|
||||
elif self.weighting_scheme == "mode":
|
||||
t = 1 - t - self.mode_scale * (torch.cos(math.pi * t / 2) ** 2 - 1 + t)
|
||||
return t
|
||||
|
||||
def sigma_to_t(self, sigma, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
||||
def t_to_sigma(self, t, **kwargs):
|
||||
indices = t.long()
|
||||
indices[indices >= self.num_timesteps] = self.num_timesteps - 1
|
||||
timesteps = self.timesteps.to(t)[indices]
|
||||
step_indices = [(self.timesteps.to(t) == t).nonzero().item() for t in timesteps]
|
||||
sigma = self.sigmas[step_indices].flatten().to(t)
|
||||
return sigma
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('NOISE_SCHEDULER',
|
||||
__class__.__name__,
|
||||
FlowMatchSigmaScheduler.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
from scepter.modules.utils.config import Config
|
||||
cfg = Config(cfg_dict={
|
||||
"NAME": "FlowMatchShiftScheduler",
|
||||
"SHIFT": 1.15
|
||||
}, load=False)
|
||||
|
||||
scheduler = NOISE_SCHEDULERS.build(cfg)
|
||||
@@ -0,0 +1,10 @@
|
||||
import torch
|
||||
|
||||
def _i(tensor, t, x):
|
||||
"""
|
||||
Index tensor using t and format the output according to x.
|
||||
"""
|
||||
shape = (x.size(0),) + (1,) * (x.ndim - 1)
|
||||
if isinstance(t, torch.Tensor):
|
||||
t = t.to(tensor.device)
|
||||
return tensor[t].view(shape).to(x.device)
|
||||
@@ -5,3 +5,4 @@ from scepter.modules.model.embedder.embedder import (
|
||||
ConcatTimestepEmbedderND, FrozenCLIPEmbedder, FrozenCLIPEmbedder2,
|
||||
FrozenOpenCLIPEmbedder, FrozenOpenCLIPEmbedder2, GeneralConditioner,
|
||||
IPAdapterPlusEmbedder, RefCrossEmbedder, SD3TextEmbedder, T5EmbedderHF)
|
||||
from scepter.modules.model.embedder.flux_embedder import HFEmbedder
|
||||
|
||||
@@ -0,0 +1,163 @@
|
||||
import torch
|
||||
from scepter.modules.model.embedder.base_embedder import BaseEmbedder
|
||||
from scepter.modules.model.registry import EMBEDDERS
|
||||
from scepter.modules.model.tokenizer.tokenizer_component import whitespace_clean, basic_clean, canonicalize
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
import transformers
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
|
||||
@EMBEDDERS.register_class()
|
||||
class HFEmbedder(BaseEmbedder):
|
||||
para_dict = {
|
||||
"HF_MODEL_CLS": {
|
||||
"value": None,
|
||||
"description": "huggingface cls in transfomer"
|
||||
},
|
||||
"MODEL_PATH": {
|
||||
"value": None,
|
||||
"description": "model folder path"
|
||||
},
|
||||
"HF_TOKENIZER_CLS": {
|
||||
"value": None,
|
||||
"description": "huggingface cls in transfomer"
|
||||
},
|
||||
|
||||
"TOKENIZER_PATH": {
|
||||
"value": None,
|
||||
"description": "tokenizer folder path"
|
||||
},
|
||||
"MAX_LENGTH": {
|
||||
"value": 77,
|
||||
"description": "max length of input"
|
||||
},
|
||||
"OUTPUT_KEY": {
|
||||
"value": "last_hidden_state",
|
||||
"description": "output key"
|
||||
},
|
||||
"D_TYPE": {
|
||||
"value": "float",
|
||||
"description": "dtype"
|
||||
},
|
||||
"BATCH_INFER": {
|
||||
"value": False,
|
||||
"description": "batch infer"
|
||||
}
|
||||
}
|
||||
para_dict.update(BaseEmbedder.para_dict)
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
hf_model_cls = cfg.get('HF_MODEL_CLS', None)
|
||||
model_path = cfg.get("MODEL_PATH", None)
|
||||
hf_tokenizer_cls = cfg.get('HF_TOKENIZER_CLS', None)
|
||||
tokenizer_path = cfg.get('TOKENIZER_PATH', None)
|
||||
self.max_length = cfg.get('MAX_LENGTH', 77)
|
||||
self.output_key = cfg.get("OUTPUT_KEY", "last_hidden_state")
|
||||
self.d_type = cfg.get("D_TYPE", "float")
|
||||
self.clean = cfg.get("CLEAN", "whitespace")
|
||||
self.batch_infer = cfg.get("BATCH_INFER", False)
|
||||
torch_dtype = getattr(torch, self.d_type)
|
||||
|
||||
assert hf_model_cls is not None and hf_tokenizer_cls is not None
|
||||
assert model_path is not None and tokenizer_path is not None
|
||||
|
||||
with FS.get_dir_to_local_dir(tokenizer_path, wait_finish=True) as local_path:
|
||||
self.tokenizer = getattr(transformers, hf_tokenizer_cls).from_pretrained(local_path,
|
||||
max_length = self.max_length,
|
||||
torch_dtype = torch_dtype)
|
||||
|
||||
with FS.get_dir_to_local_dir(model_path, wait_finish=True) as local_path:
|
||||
self.hf_module = getattr(transformers, hf_model_cls).from_pretrained(local_path, torch_dtype = torch_dtype)
|
||||
|
||||
|
||||
self.hf_module = self.hf_module.eval().requires_grad_(False)
|
||||
|
||||
def forward(self, text: list[str], return_mask = False):
|
||||
batch_encoding = self.tokenizer(
|
||||
text,
|
||||
truncation=True,
|
||||
max_length=self.max_length,
|
||||
return_length=False,
|
||||
return_overflowing_tokens=False,
|
||||
padding="max_length",
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
outputs = self.hf_module(
|
||||
input_ids=batch_encoding["input_ids"].to(self.hf_module.device),
|
||||
attention_mask=None,
|
||||
output_hidden_states=False,
|
||||
)
|
||||
if return_mask:
|
||||
return outputs[self.output_key], batch_encoding['attention_mask'].to(self.hf_module.device)
|
||||
else:
|
||||
return outputs[self.output_key], None
|
||||
|
||||
def encode(self, text, return_mask = False):
|
||||
if isinstance(text, str):
|
||||
text = [text]
|
||||
if self.clean:
|
||||
text = [self._clean(u) for u in text]
|
||||
if not self.batch_infer:
|
||||
cont, mask = [], []
|
||||
for tt in text:
|
||||
one_cont, one_mask = self([tt], return_mask=return_mask)
|
||||
cont.append(one_cont)
|
||||
mask.append(one_mask)
|
||||
if return_mask:
|
||||
return torch.cat(cont, dim=0), torch.cat(mask, dim=0)
|
||||
else:
|
||||
return torch.cat(cont, dim=0)
|
||||
else:
|
||||
ret_data = self(text, return_mask = return_mask)
|
||||
if return_mask:
|
||||
return ret_data
|
||||
else:
|
||||
return ret_data[0]
|
||||
|
||||
|
||||
def _clean(self, text):
|
||||
if self.clean == 'whitespace':
|
||||
text = whitespace_clean(basic_clean(text))
|
||||
elif self.clean == 'lower':
|
||||
text = whitespace_clean(basic_clean(text)).lower()
|
||||
elif self.clean == 'canonicalize':
|
||||
text = canonicalize(basic_clean(text))
|
||||
return text
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('EMBEDDER',
|
||||
__class__.__name__,
|
||||
HFEmbedder.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@EMBEDDERS.register_class()
|
||||
class T5PlusClipFluxEmbedder(BaseEmbedder):
|
||||
"""
|
||||
Uses the OpenCLIP transformer encoder for text
|
||||
"""
|
||||
para_dict = {
|
||||
'T5_MODEL': {},
|
||||
'CLIP_MODEL': {}
|
||||
}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.t5_model = EMBEDDERS.build(cfg.T5_MODEL, logger=logger)
|
||||
self.clip_model = EMBEDDERS.build(cfg.CLIP_MODEL, logger=logger)
|
||||
|
||||
def encode(self, text):
|
||||
t5_embeds = self.t5_model.encode(text, return_mask = False)
|
||||
clip_embeds = self.clip_model.encode(text, return_mask = False)
|
||||
# change embedding strategy here
|
||||
return {
|
||||
'context': t5_embeds,
|
||||
'y': clip_embeds,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('EMBEDDER',
|
||||
__class__.__name__,
|
||||
T5PlusClipFluxEmbedder.para_dict,
|
||||
set_name=True)
|
||||
@@ -5,4 +5,5 @@ from scepter.modules.model.network.classifier import Classifier
|
||||
from scepter.modules.model.network.diffusion import (diffusion, schedules,
|
||||
solvers)
|
||||
from scepter.modules.model.network.ldm import (ldm, ldm_edit, ldm_pixart,
|
||||
ldm_sce, ldm_sd3, ldm_xl)
|
||||
ldm_sce, ldm_sd3, ldm_xl,
|
||||
ldm_flux)
|
||||
@@ -1,18 +1,20 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import numbers
|
||||
import re
|
||||
from collections import OrderedDict
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from scepter.modules.model.network.train_module import TrainModule
|
||||
from scepter.modules.model.registry import BACKBONES, LOSSES, MODELS
|
||||
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 einops import repeat
|
||||
import math
|
||||
import torch.nn.functional as F
|
||||
class DiagonalGaussianDistribution(object):
|
||||
def __init__(self, mean, logvar, deterministic=False):
|
||||
self.mean = mean
|
||||
@@ -77,6 +79,11 @@ class AutoencoderKL(TrainModule):
|
||||
'value': 16,
|
||||
'description': ''
|
||||
},
|
||||
'SCALE_FACTOR': {
|
||||
'value': None,
|
||||
'description':
|
||||
'if is not None, will used to scale the latent space.'
|
||||
},
|
||||
}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
@@ -89,6 +96,7 @@ class AutoencoderKL(TrainModule):
|
||||
self.ignore_keys = self.cfg.get('IGNORE_KEYS', [])
|
||||
self.batch_size = self.cfg.get('BATCH_SIZE', 16)
|
||||
self.use_conv = self.cfg.get('USE_CONV', True)
|
||||
self.scale_factor = self.cfg.get('SCALE_FACTOR', None)
|
||||
|
||||
self.construct_network()
|
||||
self.init_network()
|
||||
@@ -169,6 +177,7 @@ class AutoencoderKL(TrainModule):
|
||||
return z
|
||||
|
||||
def _encode(self, x, return_mom=False):
|
||||
|
||||
h = self.encoder(x)
|
||||
moments = self.conv1(h)
|
||||
if return_mom:
|
||||
@@ -176,9 +185,15 @@ class AutoencoderKL(TrainModule):
|
||||
mean, logvar = torch.chunk(moments, 2, dim=1)
|
||||
posterior = DiagonalGaussianDistribution(mean, logvar)
|
||||
z = posterior.sample()
|
||||
if self.scale_factor is not None and isinstance(
|
||||
self.scale_factor, numbers.Number):
|
||||
z = self.scale_factor * z
|
||||
return z
|
||||
|
||||
def _decode(self, z):
|
||||
if self.scale_factor is not None and isinstance(
|
||||
self.scale_factor, numbers.Number):
|
||||
z = z / self.scale_factor
|
||||
z = self.conv2(z)
|
||||
dec = self.decoder(z)
|
||||
return dec
|
||||
@@ -268,13 +283,274 @@ class AutoencoderKL(TrainModule):
|
||||
AutoencoderKL.para_dict,
|
||||
set_name=True)
|
||||
|
||||
def timestep_embedding(timesteps, dim, max_period=10000, repeat_only=False):
|
||||
"""
|
||||
Create sinusoidal timestep embeddings.
|
||||
:param timesteps: a 1-D Tensor of N indices, one per batch element.
|
||||
These may be fractional.
|
||||
:param dim: the dimension of the output.
|
||||
:param max_period: controls the minimum frequency of the embeddings.
|
||||
:return: an [N x dim] Tensor of positional embeddings.
|
||||
"""
|
||||
if not repeat_only:
|
||||
half = dim // 2
|
||||
freqs = torch.exp(
|
||||
-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half
|
||||
).to(device=timesteps.device)
|
||||
args = torch.mm(timesteps.float().unsqueeze(1), freqs.unsqueeze(0)).view(timesteps.shape[0], len(freqs))
|
||||
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
||||
if dim % 2:
|
||||
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
||||
else:
|
||||
embedding = repeat(timesteps, 'b -> b d', d=dim)
|
||||
return embedding
|
||||
|
||||
if __name__ == '__main__':
|
||||
import argparse
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.logger import get_logger
|
||||
std_logger = get_logger(name='scepter')
|
||||
parser = argparse.ArgumentParser(description='Argparser for Scepter:\n')
|
||||
cfg = Config(load=True, parser_ins=parser)
|
||||
model = AutoencoderKL(cfg, logger=std_logger)
|
||||
model.load_pretrained_model(cfg.PRETRAINED_MODEL)
|
||||
@MODELS.register_class()
|
||||
class AutoencoderKLFlux(TrainModule):
|
||||
para_dict = {
|
||||
"ENCODER": {},
|
||||
"DECODER": {},
|
||||
"LOSS": {},
|
||||
"EMBED_DIM": {
|
||||
"value": 4,
|
||||
"description": ""
|
||||
},
|
||||
"PRETRAINED_MODEL": {
|
||||
"value": None,
|
||||
"description": ""
|
||||
},
|
||||
"IGNORE_KEYS": {
|
||||
"value": [],
|
||||
"description": ""
|
||||
},
|
||||
"BATCH_SIZE": {
|
||||
"value": 16,
|
||||
"description": ""
|
||||
},
|
||||
}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.encoder_cfg = self.cfg.ENCODER
|
||||
self.decoder_cfg = self.cfg.DECODER
|
||||
self.loss_cfg = self.cfg.get("LOSS", None)
|
||||
self.embed_dim = self.cfg.get("EMBED_DIM", 4)
|
||||
self.pretrained_model = self.cfg.get("PRETRAINED_MODEL", None)
|
||||
#
|
||||
self.ignore_keys = self.cfg.get("IGNORE_KEYS", [])
|
||||
self.batch_size = self.cfg.get("BATCH_SIZE", 16)
|
||||
self.resize_nx = self.cfg.get("RESIZE_NX", 1)
|
||||
self.use_rembed = self.cfg.get("USE_REMBED", True)
|
||||
self.use_conv = self.cfg.get('USE_CONV', True)
|
||||
self.scale_factor = self.cfg.get('SCALE_FACTOR', None)
|
||||
self.shift_factor = self.cfg.get('SHIFT_FACTOR', None)
|
||||
self.construct_network()
|
||||
|
||||
def construct_network(self):
|
||||
z_channels = self.encoder_cfg.Z_CHANNELS
|
||||
self.encoder = BACKBONES.build(self.encoder_cfg, logger=self.logger)
|
||||
self.decoder = BACKBONES.build(self.decoder_cfg, logger=self.logger)
|
||||
self.conv1 = torch.nn.Conv2d(2 * z_channels, 2 * self.embed_dim, 1) if self.use_conv else torch.nn.Identity()
|
||||
self.conv2 = torch.nn.Conv2d(self.embed_dim, z_channels,1) if self.use_conv else torch.nn.Identity()
|
||||
|
||||
# freeze encoder
|
||||
for param in self.encoder.parameters():
|
||||
param.requires_grad = False
|
||||
for param in self.conv1.parameters():
|
||||
param.requires_grad = False
|
||||
for param in self.conv2.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def load_pretrained_model(self, pretrained_model):
|
||||
if pretrained_model is not None:
|
||||
with FS.get_from(pretrained_model, wait_finish=True) as local_model:
|
||||
self.init_from_ckpt(local_model)
|
||||
def init_from_ckpt(self, path, ignore_keys=list()):
|
||||
if path.find('.safetensors') > -1:
|
||||
from safetensors import safe_open
|
||||
sd = OrderedDict()
|
||||
with safe_open(path, framework="pt", device='cpu') as f:
|
||||
for k in f.keys():
|
||||
sd[k] = f.get_tensor(k)
|
||||
else:
|
||||
sd = torch.load(path, map_location="cpu")
|
||||
if path.find('.pt') > -1 and 'state_dict' in sd:
|
||||
sd = sd['state_dict']
|
||||
elif path.find('.ckpt') > -1 and 'state_dict' in sd:
|
||||
sd = sd['state_dict']
|
||||
elif path.find('.pth') > -1 and 'model' in sd:
|
||||
sd = sd['model']
|
||||
|
||||
new_sd = OrderedDict()
|
||||
for k, v in sd.items():
|
||||
ignored = False
|
||||
for ik in ignore_keys:
|
||||
if ik in k:
|
||||
if we.rank == 0:
|
||||
self.logger.info("ignore key {} from state_dict.".format(k))
|
||||
ignored = True
|
||||
break
|
||||
k = k.replace("post_quant_conv", "conv2") if "post_quant_conv" in k else k
|
||||
k = k.replace("quant_conv", "conv1") if "quant_conv" in k else k
|
||||
if not ignored:
|
||||
new_sd[k] = v
|
||||
|
||||
missing, unexpected = self.load_state_dict(new_sd, strict=False)
|
||||
if we.rank == 0:
|
||||
self.logger.info(f"Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys")
|
||||
if len(missing) > 0:
|
||||
self.logger.info(f"Missing Keys:\n {missing}")
|
||||
if len(unexpected) > 0:
|
||||
self.logger.info(f"\nUnexpected Keys:\n {unexpected}")
|
||||
|
||||
@torch.no_grad()
|
||||
def encode(self, x, sample_posterior = True):
|
||||
h = self.encoder(x)
|
||||
moments = self.conv1(h)
|
||||
mean, logvar = torch.chunk(moments, 2, dim=1)
|
||||
posterior = DiagonalGaussianDistribution(mean, logvar)
|
||||
if sample_posterior:
|
||||
z = posterior.sample()
|
||||
else:
|
||||
z = posterior.mode()
|
||||
if self.shift_factor is not None and isinstance(self.shift_factor, numbers.Number):
|
||||
z = z - self.shift_factor
|
||||
if self.scale_factor is not None and isinstance(self.scale_factor, numbers.Number):
|
||||
z = self.scale_factor * z
|
||||
return z, posterior
|
||||
|
||||
def decode(self, z, **kwargs):
|
||||
b, c, h, w = z.size()
|
||||
if kwargs.get('resize_nx', None) is not None and self.use_rembed:
|
||||
resize_nx = kwargs['resize_nx']
|
||||
if not torch.is_tensor(resize_nx):
|
||||
resize_nx = torch.full((b,), resize_nx, device=we.device_id, dtype=z.dtype)
|
||||
rembed = timestep_embedding(resize_nx, dim=self.decoder_cfg.CH_MULT[-1] * self.decoder_cfg.CH)
|
||||
else:
|
||||
rembed = None
|
||||
if self.scale_factor is not None and isinstance(self.scale_factor, numbers.Number):
|
||||
z = z / self.scale_factor
|
||||
if self.shift_factor is not None and isinstance(self.shift_factor, numbers.Number):
|
||||
z = z + self.shift_factor
|
||||
z = self.conv2(z)
|
||||
# add grad;
|
||||
if rembed is not None:
|
||||
dec = self.decoder(z, rembed)
|
||||
else:
|
||||
dec = self.decoder(z)
|
||||
return dec
|
||||
|
||||
def share_forward(self, image=None, sample_posterior=True, **kwargs):
|
||||
# rembed: resize embedding
|
||||
if image is not None:
|
||||
z, posterior = self.encode(image, sample_posterior = sample_posterior)
|
||||
else:
|
||||
latent = kwargs.pop("latent", None)
|
||||
assert latent is not None
|
||||
z = latent
|
||||
posterior = None
|
||||
if self.shift_factor is not None and isinstance(self.scale_factor, numbers.Number):
|
||||
z = z - self.shift_factor
|
||||
if self.scale_factor is not None and isinstance(self.scale_factor, numbers.Number):
|
||||
z = self.scale_factor * z
|
||||
dec = self.decode(z, **kwargs)
|
||||
return dec, posterior
|
||||
|
||||
def forward(self, **kwargs):
|
||||
if self.training:
|
||||
ret = self.forward_train(**kwargs)
|
||||
else:
|
||||
ret = self.forward_test(**kwargs)
|
||||
return ret
|
||||
|
||||
def forward_train(self,
|
||||
image=None,
|
||||
gt_image=None,
|
||||
sample_posterior=True,
|
||||
optimizer_idx=0,
|
||||
global_step=0,
|
||||
**kwargs):
|
||||
|
||||
if gt_image is None:
|
||||
gt_image = copy.deepcopy(image)
|
||||
reconstructions, posterior = self.share_forward(image, sample_posterior, **kwargs)
|
||||
|
||||
ret = {}
|
||||
if optimizer_idx == 0:
|
||||
# train encoder+decoder+logvar
|
||||
aeloss, log_dict_ae = self.loss(gt_image, reconstructions, posterior, optimizer_idx,
|
||||
global_step, last_layer=self.get_last_layer(), split="train")
|
||||
# self.logger.info(f"aeloss: {aeloss.detach().cpu().item()}, ")
|
||||
ret["loss"] = aeloss
|
||||
ret.update(log_dict_ae)
|
||||
|
||||
if optimizer_idx == 1:
|
||||
# train the discriminator
|
||||
discloss, log_dict_disc = self.loss(gt_image, reconstructions, posterior, optimizer_idx,
|
||||
global_step, last_layer=self.get_last_layer(), split="train")
|
||||
# self.logger.info(f"discloss: {discloss.detach().cpu().item()}, ")
|
||||
ret["loss"] = discloss
|
||||
ret.update(log_dict_disc)
|
||||
|
||||
return ret
|
||||
|
||||
@torch.no_grad()
|
||||
def forward_test(self,
|
||||
image=None,
|
||||
gt_image=None,
|
||||
sample_posterior=True,
|
||||
**kwargs):
|
||||
resize_nx_ = 1
|
||||
if image is not None:
|
||||
b, c, h, w = image.size()
|
||||
if kwargs.get('resize_ex'):
|
||||
resize_nx_ = kwargs.pop('resize_ex')
|
||||
image = F.interpolate(image, (int(float(h) / resize_nx_), int(float(w) / resize_nx_)), mode='bicubic')
|
||||
image = F.interpolate(image, (h, w), mode='bicubic')
|
||||
kwargs["resize_nx"] = resize_nx_
|
||||
elif kwargs.get('resize_nx', None) is not None:
|
||||
resize_nx_ = kwargs['resize_nx']
|
||||
|
||||
if gt_image is None:
|
||||
if image is not None:
|
||||
gt_image = copy.deepcopy(image)
|
||||
|
||||
# kwargs["resize_nx"] = resize_nx_
|
||||
reconstructions, posterior = self.share_forward(image, sample_posterior, **kwargs)
|
||||
reconstructions = torch.clamp((reconstructions + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
if gt_image is not None:
|
||||
gt_image = torch.clamp((gt_image + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
else:
|
||||
gt_image = [None for _ in range(reconstructions.shape[0])]
|
||||
if image is not None:
|
||||
lr_image = torch.clamp((image + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
else:
|
||||
lr_image = [None for _ in range(reconstructions.shape[0])]
|
||||
|
||||
ret = list()
|
||||
if torch.is_tensor(resize_nx_):
|
||||
resize_nx = [nx.item() for nx, in zip(resize_nx_.cpu())]
|
||||
else:
|
||||
resize_nx = [resize_nx_ for _ in range(reconstructions.size(0))]
|
||||
|
||||
for img, ori, gt_img, nx in zip(reconstructions, lr_image, gt_image, resize_nx):
|
||||
ret.append({
|
||||
"prompt": "",
|
||||
"n_prompt": "",
|
||||
"image": img,
|
||||
"lr_image": ori,
|
||||
"gt_image": gt_img,
|
||||
"resize_nx": nx
|
||||
})
|
||||
|
||||
return ret
|
||||
|
||||
def get_last_layer(self):
|
||||
if hasattr(self.decoder, 'conv_out'):
|
||||
return self.decoder.conv_out.weight
|
||||
else:
|
||||
return self.decoder.head[-1].weight
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml("MODEL", __class__.__name__, AutoencoderKLFlux.para_dict, set_name=True)
|
||||
|
||||
@@ -10,7 +10,7 @@ import torch
|
||||
from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion
|
||||
from scepter.modules.model.network.diffusion.schedules import noise_schedule
|
||||
from scepter.modules.model.network.train_module import TrainModule
|
||||
from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, LOSSES,
|
||||
from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, LOSSES, DIFFUSIONS,
|
||||
MODELS, TOKENIZERS)
|
||||
from scepter.modules.model.utils.basic_utils import count_params, default
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
@@ -112,13 +112,18 @@ class LatentDiffusion(TrainModule):
|
||||
if self.zero_terminal_snr:
|
||||
assert self.parameterization == 'v', 'Now zero_terminal_snr only support v-prediction mode.'
|
||||
|
||||
self.sigmas = noise_schedule(schedule=self.schedule_args.pop('name'),
|
||||
n=self.num_timesteps,
|
||||
zero_terminal_snr=self.zero_terminal_snr,
|
||||
**self.schedule_args)
|
||||
|
||||
self.diffusion = GaussianDiffusion(
|
||||
sigmas=self.sigmas, prediction_type=self.parameterization)
|
||||
diffusion_cfg = self.cfg.get("DIFFUSION", None)
|
||||
if diffusion_cfg is not None:
|
||||
if self.cfg.have("WORK_DIR"):
|
||||
diffusion_cfg.WORK_DIR = self.cfg.WORK_DIR
|
||||
self.diffusion = DIFFUSIONS.build(diffusion_cfg, logger=self.logger)
|
||||
else:
|
||||
self.sigmas = noise_schedule(schedule=self.schedule_args.pop('name'),
|
||||
n=self.num_timesteps,
|
||||
zero_terminal_snr=self.zero_terminal_snr,
|
||||
**self.schedule_args)
|
||||
self.diffusion = GaussianDiffusion(
|
||||
sigmas=self.sigmas, prediction_type=self.parameterization)
|
||||
|
||||
self.pretrained_model = self.cfg.get('PRETRAINED_MODEL', None)
|
||||
self.ignore_keys = self.cfg.get('IGNORE_KEYS', [])
|
||||
@@ -131,6 +136,7 @@ class LatentDiffusion(TrainModule):
|
||||
|
||||
self.scale_factor = self.cfg.get('SCALE_FACTOR', 0.18215)
|
||||
self.size_factor = self.cfg.get('SIZE_FACTOR', 8)
|
||||
self.decoder_bias = self.cfg.get("DECODER_BIAS", 0)
|
||||
self.default_n_prompt = self.cfg.get('DEFAULT_N_PROMPT', '')
|
||||
self.default_n_prompt = '' if self.default_n_prompt is None else self.default_n_prompt
|
||||
self.p_zero = self.cfg.get('P_ZERO', 0.0)
|
||||
@@ -140,6 +146,7 @@ class LatentDiffusion(TrainModule):
|
||||
if self.train_n_prompt is None:
|
||||
self.train_n_prompt = ''
|
||||
self.use_ema = self.cfg.get('USE_EMA', False)
|
||||
self.eval_ema = self.cfg.get('EVAL_EMA', False)
|
||||
self.model_ema_config = self.cfg.get('DIFFUSION_MODEL_EMA', None)
|
||||
|
||||
def construct_network(self):
|
||||
@@ -290,6 +297,7 @@ class LatentDiffusion(TrainModule):
|
||||
@torch.no_grad()
|
||||
@torch.autocast('cuda', dtype=torch.float16)
|
||||
def forward_test(self,
|
||||
image=None,
|
||||
prompt=None,
|
||||
n_prompt=None,
|
||||
sampler='ddim',
|
||||
|
||||
@@ -111,6 +111,7 @@ class LatentDiffusionEdit(LatentDiffusion):
|
||||
@torch.no_grad()
|
||||
@torch.autocast('cuda', dtype=torch.float16)
|
||||
def forward_test(self,
|
||||
image=None,
|
||||
prompt=None,
|
||||
n_prompt=None,
|
||||
sampler='ddim',
|
||||
|
||||
@@ -0,0 +1,220 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import math
|
||||
import numbers
|
||||
import random
|
||||
import torch
|
||||
from scepter.modules.model.network.ldm import LatentDiffusion
|
||||
from scepter.modules.model.registry import MODELS, BACKBONES, LOSSES, TOKENIZERS, EMBEDDERS, DIFFUSIONS
|
||||
from scepter.modules.model.utils.basic_utils import disabled_train
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.model.utils.basic_utils import count_params
|
||||
|
||||
|
||||
@MODELS.register_class()
|
||||
class LatentDiffusionFlux(LatentDiffusion):
|
||||
para_dict = LatentDiffusion.para_dict
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.guide_scale = cfg.get('GUIDE_SCALE', 3.5)
|
||||
|
||||
def init_params(self):
|
||||
self.parameterization = self.cfg.get('PARAMETERIZATION', 'rf')
|
||||
assert self.parameterization in [
|
||||
'eps', 'x0', 'v', 'rf'
|
||||
], 'currently only supporting "eps" and "x0" and "v" and "rf"'
|
||||
|
||||
diffusion_cfg = self.cfg.get("DIFFUSION", None)
|
||||
assert diffusion_cfg is not None
|
||||
if self.cfg.have("WORK_DIR"):
|
||||
diffusion_cfg.WORK_DIR = self.cfg.WORK_DIR
|
||||
self.diffusion = DIFFUSIONS.build(diffusion_cfg, logger=self.logger)
|
||||
|
||||
self.pretrained_model = self.cfg.get('PRETRAINED_MODEL', None)
|
||||
self.ignore_keys = self.cfg.get('IGNORE_KEYS', [])
|
||||
|
||||
self.model_config = self.cfg.DIFFUSION_MODEL
|
||||
self.first_stage_config = self.cfg.FIRST_STAGE_MODEL
|
||||
self.cond_stage_config = self.cfg.COND_STAGE_MODEL
|
||||
self.tokenizer_config = self.cfg.get('TOKENIZER', None)
|
||||
self.loss_config = self.cfg.get('LOSS', None)
|
||||
|
||||
self.scale_factor = self.cfg.get('SCALE_FACTOR', 0.18215)
|
||||
self.size_factor = self.cfg.get('SIZE_FACTOR', 16)
|
||||
self.default_n_prompt = self.cfg.get('DEFAULT_N_PROMPT', '')
|
||||
self.default_n_prompt = '' if self.default_n_prompt is None else self.default_n_prompt
|
||||
self.p_zero = self.cfg.get('P_ZERO', 0.0)
|
||||
self.train_n_prompt = self.cfg.get('TRAIN_N_PROMPT', '')
|
||||
if self.default_n_prompt is None:
|
||||
self.default_n_prompt = ''
|
||||
if self.train_n_prompt is None:
|
||||
self.train_n_prompt = ''
|
||||
self.use_ema = self.cfg.get('USE_EMA', False)
|
||||
self.model_ema_config = self.cfg.get('DIFFUSION_MODEL_EMA', None)
|
||||
|
||||
def construct_network(self):
|
||||
# embedding_context = torch.device("meta") if self.model_config.get("PRETRAINED_MODEL", None) else nullcontext()
|
||||
# with embedding_context:
|
||||
self.model = BACKBONES.build(self.model_config, logger=self.logger).to(torch.bfloat16)
|
||||
self.logger.info('all parameters:{}'.format(count_params(self.model)))
|
||||
if self.use_ema:
|
||||
if self.model_ema_config:
|
||||
self.model_ema = BACKBONES.build(self.model_ema_config,
|
||||
logger=self.logger)
|
||||
else:
|
||||
self.model_ema = copy.deepcopy(self.model)
|
||||
self.model_ema = self.model_ema.eval()
|
||||
for param in self.model_ema.parameters():
|
||||
param.requires_grad = False
|
||||
if self.loss_config:
|
||||
self.loss = LOSSES.build(self.loss_config, logger=self.logger)
|
||||
if self.tokenizer_config is not None:
|
||||
self.tokenizer = TOKENIZERS.build(self.tokenizer_config,
|
||||
logger=self.logger)
|
||||
if self.first_stage_config:
|
||||
self.first_stage_model = MODELS.build(self.first_stage_config,
|
||||
logger=self.logger)
|
||||
self.first_stage_model = self.first_stage_model.eval()
|
||||
self.first_stage_model.train = disabled_train
|
||||
for param in self.first_stage_model.parameters():
|
||||
param.requires_grad = False
|
||||
else:
|
||||
self.first_stage_model = None
|
||||
if self.tokenizer_config is not None:
|
||||
self.cond_stage_config.KWARGS = {
|
||||
'vocab_size': self.tokenizer.vocab_size
|
||||
}
|
||||
if self.cond_stage_config == '__is_unconditional__':
|
||||
print(
|
||||
f'Training {self.__class__.__name__} as an unconditional model.'
|
||||
)
|
||||
self.cond_stage_model = None
|
||||
else:
|
||||
model = EMBEDDERS.build(self.cond_stage_config, logger=self.logger)
|
||||
self.cond_stage_model = model.eval().requires_grad_(False)
|
||||
self.cond_stage_model.train = disabled_train
|
||||
|
||||
def noise_sample(self, num_samples, h, w, seed, dtype = torch.bfloat16):
|
||||
noise = torch.randn(
|
||||
num_samples,
|
||||
16,
|
||||
# allow for packing
|
||||
2 * math.ceil(h / 16),
|
||||
2 * math.ceil(w / 16),
|
||||
device=we.device_id,
|
||||
dtype=dtype,
|
||||
generator=torch.Generator(device=we.device_id).manual_seed(seed),
|
||||
)
|
||||
return noise
|
||||
|
||||
def forward_train(self, image=None, noise=None, prompt=None, **kwargs):
|
||||
x_start = self.encode_first_stage(image, **kwargs)
|
||||
if prompt and self.cond_stage_model:
|
||||
ctx = getattr(self.cond_stage_model, 'encode')(prompt)
|
||||
else:
|
||||
assert False
|
||||
if 'index' in kwargs:
|
||||
kwargs.pop('index')
|
||||
guide_scale = self.guide_scale
|
||||
if guide_scale is not None:
|
||||
guide_scale = torch.full((x_start.shape[0],), guide_scale, device=x_start.device, dtype=x_start.dtype)
|
||||
else:
|
||||
guide_scale = None
|
||||
loss = self.diffusion.loss(x_0=x_start,
|
||||
model=self.model,
|
||||
model_kwargs={"cond": ctx, "guidance": guide_scale},
|
||||
noise=noise,
|
||||
**kwargs)
|
||||
loss = loss.mean()
|
||||
ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
|
||||
return ret
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def forward_test(self,
|
||||
image=None,
|
||||
prompt=None,
|
||||
sampler='flow_eluer',
|
||||
sample_steps=20,
|
||||
seed=2023,
|
||||
guide_scale=4.5,
|
||||
guide_rescale=0.0,
|
||||
show_process=False,
|
||||
**kwargs):
|
||||
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
||||
if isinstance(prompt, str):
|
||||
prompt = [prompt]
|
||||
assert isinstance(prompt, list)
|
||||
num_samples = len(prompt)
|
||||
if prompt and self.cond_stage_model:
|
||||
ctx = getattr(self.cond_stage_model, 'encode')(prompt)
|
||||
else:
|
||||
assert False
|
||||
if 'index' in kwargs:
|
||||
kwargs.pop('index')
|
||||
image_size = None
|
||||
if 'meta' in kwargs:
|
||||
meta = kwargs.pop('meta')
|
||||
if 'image_size' in meta:
|
||||
h = int(meta['image_size'][0][0])
|
||||
w = int(meta['image_size'][1][0])
|
||||
image_size = [h, w]
|
||||
if 'image_size' in kwargs:
|
||||
image_size = kwargs.pop('image_size')
|
||||
if isinstance(image_size, numbers.Number):
|
||||
image_size = [image_size, image_size]
|
||||
if image_size is None:
|
||||
image_size = [1024, 1024]
|
||||
height, width = image_size
|
||||
noise = self.noise_sample(
|
||||
num_samples,
|
||||
height,
|
||||
width,
|
||||
seed
|
||||
)
|
||||
guide_scale = guide_scale or self.guide_scale
|
||||
if guide_scale is not None:
|
||||
guide_scale = torch.full((noise.shape[0],), guide_scale, device=noise.device, dtype=noise.dtype)
|
||||
else:
|
||||
guide_scale = None
|
||||
# UNet use input n_prompt
|
||||
samples = self.diffusion.sample(
|
||||
noise=noise,
|
||||
sampler=sampler,
|
||||
model=self.model,
|
||||
model_kwargs= {"cond": ctx, "guidance": guide_scale},
|
||||
steps=sample_steps,
|
||||
show_progress=True,
|
||||
guide_scale = guide_scale,
|
||||
return_intermediate=None,
|
||||
**kwargs).float()
|
||||
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
|
||||
x_samples = self.decode_first_stage(samples).float()
|
||||
x_samples = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
outputs = list()
|
||||
|
||||
for i, (p, img) in enumerate(zip(prompt, x_samples)):
|
||||
one_tup = {'prompt': str(p), 'n_prompt': '', 'image': img}
|
||||
outputs.append(one_tup)
|
||||
|
||||
return outputs
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('MODEL',
|
||||
__class__.__name__,
|
||||
LatentDiffusionFlux.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_first_stage(self, x, **kwargs):
|
||||
z = self.first_stage_model.encode(x)
|
||||
if isinstance(z, (tuple, list)):
|
||||
z = z[0]
|
||||
return z
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_first_stage(self, z):
|
||||
return self.first_stage_model.decode(z)
|
||||
@@ -134,6 +134,7 @@ class LatentDiffusionPixart(LatentDiffusion):
|
||||
|
||||
@torch.no_grad()
|
||||
def forward_test(self,
|
||||
image=None,
|
||||
prompt=None,
|
||||
label=None,
|
||||
sampler='ddim',
|
||||
|
||||
@@ -136,6 +136,7 @@ class LatentDiffusionSD3(LatentDiffusion):
|
||||
|
||||
@torch.no_grad()
|
||||
def forward_test(self,
|
||||
image=None,
|
||||
prompt=None,
|
||||
sampler='ddim',
|
||||
sample_steps=20,
|
||||
|
||||
@@ -26,6 +26,27 @@ def build_model(cfg, registry, logger=None, *args, **kwargs):
|
||||
model.load_pretrained_model(pretrain_cfg)
|
||||
return model
|
||||
|
||||
def build_diffusion(cfg, registry, logger=None, *args, **kwargs):
|
||||
""" After build model, load pretrained model if exists key `pretrain`.
|
||||
|
||||
pretrain (str, dict): Describes how to load pretrained model.
|
||||
str, treat pretrain as model path;
|
||||
dict: should contains key `path`, and other parameters token by function load_pretrained();
|
||||
"""
|
||||
if not isinstance(cfg, Config):
|
||||
raise TypeError(f'Config must be type dict, got {type(cfg)}')
|
||||
return build_from_config(cfg, registry, logger=logger, *args, **kwargs)
|
||||
|
||||
def build_scheduler(cfg, registry, logger=None, *args, **kwargs):
|
||||
if not isinstance(cfg, Config):
|
||||
raise TypeError(f'Config must be type dict, got {type(cfg)}')
|
||||
return build_from_config(cfg, registry, logger=logger, *args, **kwargs)
|
||||
|
||||
def build_diffusion_sampler(cfg, registry, logger=None, *args, **kwargs):
|
||||
if not isinstance(cfg, Config):
|
||||
raise TypeError(f'Config must be type dict, got {type(cfg)}')
|
||||
return build_from_config(cfg, registry, logger=logger, *args, **kwargs)
|
||||
|
||||
|
||||
MODELS = Registry('MODELS', build_func=build_model)
|
||||
TOKENIZERS = Registry('TOKENIZER', build_func=build_model)
|
||||
@@ -37,3 +58,9 @@ BRICKS = Registry('BRICKS', build_func=build_model)
|
||||
STEMS = BRICKS
|
||||
LOSSES = Registry('LOSSES', build_func=build_model)
|
||||
TUNERS = Registry('TUNERS', build_func=build_model)
|
||||
|
||||
# reigister cls for diffusion.
|
||||
|
||||
DIFFUSIONS = Registry('DIFFUSIONS', build_func=build_diffusion)
|
||||
NOISE_SCHEDULERS = Registry('NOISE_SCHEDULERS', build_func=build_diffusion)
|
||||
DIFFUSION_SAMPLERS = Registry('DIFFUSION_SAMPLERS', build_func=build_diffusion_sampler)
|
||||
|
||||
@@ -3,4 +3,5 @@
|
||||
|
||||
from scepter.modules.opt.lr_schedulers.define_schedulers import LinoPolyLR
|
||||
from scepter.modules.opt.lr_schedulers.official_schedulers import * # noqa
|
||||
from scepter.modules.opt.lr_schedulers.warmup import WarmupToConstantLR
|
||||
from scepter.modules.opt.lr_schedulers.warmup import (StepAnnealingLR,
|
||||
WarmupToConstantLR)
|
||||
|
||||
@@ -33,7 +33,6 @@ class WarmupToConstantLR(BaseScheduler):
|
||||
WarmupToConstantLR.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
class AnnealingLR(_LRScheduler):
|
||||
def __init__(self,
|
||||
optimizer,
|
||||
|
||||
@@ -18,7 +18,10 @@ from scepter.modules.solver.hooks import HOOKS
|
||||
from scepter.modules.utils.config import Config, dict_to_yaml
|
||||
from scepter.modules.utils.data import transfer_data_to_cuda
|
||||
from scepter.modules.utils.directory import get_relative_folder, osp_path
|
||||
from scepter.modules.utils.distribute import dist, gather_data, we
|
||||
from scepter.modules.utils.distribute import (
|
||||
dist, gather_data, we, all_reduce,
|
||||
_serialize_to_tensor, broadcast, _unserialize_from_tensor,
|
||||
all_reduce, barrier)
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.modules.utils.logger import get_logger, init_logger
|
||||
from scepter.modules.utils.probe import (ProbeData, merge_gathered_probe,
|
||||
@@ -184,6 +187,20 @@ try:
|
||||
except Exception as e:
|
||||
warnings.warn(f'{e}')
|
||||
|
||||
def async_str(text):
|
||||
broadcast_size = torch.zeros(1, dtype=torch.long).to(we.device_id)
|
||||
if we.rank == 0:
|
||||
text_tensor = _serialize_to_tensor(text).to(we.device_id)
|
||||
broadcast_size[0] = len(text_tensor)
|
||||
broadcast(broadcast_size, src=0)
|
||||
broadcast(text_tensor, src=0)
|
||||
else:
|
||||
broadcast(broadcast_size, src=0)
|
||||
text_tensor = torch.empty((broadcast_size[0],), dtype=torch.uint8).to(we.device_id)
|
||||
broadcast(text_tensor, src=0)
|
||||
text = _unserialize_from_tensor(text_tensor)
|
||||
return text
|
||||
|
||||
|
||||
class BaseSolver(object, metaclass=ABCMeta):
|
||||
""" Base Solver.
|
||||
@@ -198,18 +215,12 @@ class BaseSolver(object, metaclass=ABCMeta):
|
||||
'description': 'The precision for train process.'
|
||||
},
|
||||
'FILE_SYSTEM': {},
|
||||
'ACCU_STEP': {
|
||||
'value':
|
||||
1,
|
||||
'description':
|
||||
'When use ddp, the grad accumulate steps for each process.'
|
||||
},
|
||||
'RESUME_FROM': {
|
||||
'value': '',
|
||||
'description': 'Resume from some state of training!'
|
||||
},
|
||||
'MAX_EPOCHS': {
|
||||
'value': 10,
|
||||
'value': -1,
|
||||
'description': 'Max epochs for training.'
|
||||
},
|
||||
'NUM_FOLDS': {
|
||||
@@ -252,31 +263,31 @@ class BaseSolver(object, metaclass=ABCMeta):
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
# initialize some hyperparameters
|
||||
self.cfg = cfg
|
||||
self.logger = logger
|
||||
self.file_system = cfg.get('FILE_SYSTEM', None)
|
||||
self.work_dir: str = cfg.WORK_DIR
|
||||
self.work_dir: str = async_str(cfg.WORK_DIR)
|
||||
barrier()
|
||||
self.pl_dir = self.work_dir
|
||||
self.log_file = osp_path(self.work_dir, cfg.LOG_FILE)
|
||||
self.optimizer, self.lr_scheduler = None, None
|
||||
self.cfg = cfg
|
||||
self.logger = logger
|
||||
self.resume_from: str = cfg.RESUME_FROM
|
||||
self.max_epochs: int = cfg.MAX_EPOCHS
|
||||
self.resume_from: str = cfg.get("RESUME_FROM", None)
|
||||
self.max_epochs: int = cfg.get("MAX_EPOCHS", -1)
|
||||
self.use_pl = we.use_pl
|
||||
self.train_precision = self.cfg.get('TRAIN_PRECISION', 32)
|
||||
self._mode_set = set()
|
||||
self._mode = 'train'
|
||||
self.probe_ins = {}
|
||||
self.collect_probe_ins = {}
|
||||
self.clear_probe_ins = {}
|
||||
self._num_folds: int = 1
|
||||
if not self.use_pl:
|
||||
world_size = we.world_size
|
||||
if world_size > 1:
|
||||
self._num_folds: int = cfg.NUM_FOLDS
|
||||
self._num_folds: int = cfg.get("NUM_FOLDS", 1)
|
||||
if cfg.have('MODE'):
|
||||
self._mode_set.add(cfg.MODE)
|
||||
self._mode = cfg.MODE
|
||||
if we.is_distributed:
|
||||
self.accu_step = cfg.get('ACCU_STEP', 1)
|
||||
|
||||
self.do_step = True
|
||||
self.hooks_dict = {'train': [], 'eval': [], 'test': []}
|
||||
@@ -305,6 +316,8 @@ class BaseSolver(object, metaclass=ABCMeta):
|
||||
self._prefix = FS.get_fs_client(self.work_dir).get_prefix()
|
||||
if not FS.exists(self.work_dir):
|
||||
FS.make_dir(self.work_dir)
|
||||
assert self.cfg.have('MODEL')
|
||||
self.cfg.MODEL.WORK_DIR = self.work_dir
|
||||
self.logger.info(
|
||||
f"Parse work dir {self.work_dir}'s prefix is {self._prefix}")
|
||||
|
||||
@@ -325,6 +338,7 @@ class BaseSolver(object, metaclass=ABCMeta):
|
||||
def __setattr__(self, key, value):
|
||||
if isinstance(value, BaseModel):
|
||||
self.probe_ins[key] = value.probe_data
|
||||
self.collect_probe_ins[key] = value.collect_probe
|
||||
self.clear_probe_ins[key] = value.clear_probe
|
||||
super().__setattr__(key, value)
|
||||
|
||||
@@ -666,7 +680,7 @@ class BaseSolver(object, metaclass=ABCMeta):
|
||||
return self._iter[self._mode]
|
||||
|
||||
@property
|
||||
def probe_data(self):
|
||||
def probe_data_dict(self):
|
||||
return self._probe_data[self._mode]
|
||||
|
||||
@property
|
||||
@@ -736,8 +750,32 @@ class BaseSolver(object, metaclass=ABCMeta):
|
||||
else:
|
||||
self._dist_data[self.mode][key][k] = v
|
||||
|
||||
@property
|
||||
def collect_probe(self):
|
||||
probe_data_dict = self._probe_data[self.mode]
|
||||
for k, func in self.collect_probe_ins.items():
|
||||
for kk, vv in func().items():
|
||||
probe_data_dict[f'{k}/{kk}'] = vv
|
||||
return probe_data_dict
|
||||
|
||||
@property
|
||||
def probe_data(self): # noqa
|
||||
if hasattr(self, f'{self.mode}_pre_save_paras'):
|
||||
pre_save_paras = getattr(self, f'{self.mode}_pre_save_paras')
|
||||
save_folder = pre_save_paras['save_folder']
|
||||
save_probe_prefix = pre_save_paras['save_probe_prefix']
|
||||
step = pre_save_paras['step']
|
||||
save_image_postfix = pre_save_paras.get('save_image_postfix', 'jpg')
|
||||
save_video_postfix = pre_save_paras.get('save_video_postfix', 'mp4')
|
||||
for k, v in self.collect_probe.items():
|
||||
if save_probe_prefix is not None:
|
||||
ret_prefix = os.path.join(save_folder, save_probe_prefix)
|
||||
else:
|
||||
ret_prefix = os.path.join(save_folder, k.replace('/', '_') + f'_step_{step}')
|
||||
v.presave(prefix = ret_prefix,
|
||||
image_postfix = save_image_postfix,
|
||||
video_postfix = save_video_postfix,
|
||||
rank = we.rank)
|
||||
gather_probe_data = gather_data(self._probe_data[self.mode])
|
||||
_dist_data_list = gather_data([self._dist_data[self.mode] or {}])
|
||||
if not we.rank == 0:
|
||||
@@ -836,9 +874,10 @@ class BaseSolver(object, metaclass=ABCMeta):
|
||||
for key in keys:
|
||||
value = data_dict[key]
|
||||
if isinstance(value, torch.Tensor) and value.ndim == 0:
|
||||
if dist.is_available() and dist.is_initialized():
|
||||
if we.is_distributed:
|
||||
value = value.data.clone()
|
||||
dist.all_reduce(value.div_(dist.get_world_size()))
|
||||
all_reduce(value, group=we.data_parallel_group)
|
||||
value = value/we.data_group_world_size
|
||||
ret[key] = value
|
||||
else:
|
||||
ret[key] = value
|
||||
@@ -911,7 +950,7 @@ class BaseSolver(object, metaclass=ABCMeta):
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('solvername',
|
||||
return dict_to_yaml('SOLVER',
|
||||
__class__.__name__,
|
||||
BaseSolver.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@@ -2,11 +2,14 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import os
|
||||
import re
|
||||
import warnings
|
||||
from collections import OrderedDict, defaultdict
|
||||
from functools import partial
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.cuda.amp as amp
|
||||
import torch.nn as nn
|
||||
from scepter.modules.data.dataset import DATASETS
|
||||
from scepter.modules.opt.lr_schedulers import LR_SCHEDULERS
|
||||
from scepter.modules.opt.optimizers import OPTIMIZERS
|
||||
@@ -16,19 +19,80 @@ from scepter.modules.utils.config import Config, dict_to_yaml
|
||||
from scepter.modules.utils.data import transfer_data_to_cuda
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.probe import ProbeData
|
||||
from torch.distributed.fsdp import (BackwardPrefetch, CPUOffload,
|
||||
FullStateDictConfig,
|
||||
FullyShardedDataParallel, MixedPrecision,
|
||||
ShardingStrategy, StateDictType)
|
||||
from torch.distributed.fsdp import FullStateDictConfig
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
from torch.distributed.fsdp import (MixedPrecision, ShardingStrategy,
|
||||
StateDictType)
|
||||
from torch.distributed.fsdp.wrap import (lambda_auto_wrap_policy,
|
||||
size_based_auto_wrap_policy)
|
||||
from torch.nn.parallel import DistributedDataParallel
|
||||
from tqdm import tqdm
|
||||
|
||||
sharding_strategy_map = {
|
||||
'full_shard': ShardingStrategy.FULL_SHARD,
|
||||
'shard_grad_op': ShardingStrategy.SHARD_GRAD_OP
|
||||
'shard_grad_op': ShardingStrategy.SHARD_GRAD_OP,
|
||||
'hybrid_shard': ShardingStrategy.HYBRID_SHARD
|
||||
}
|
||||
|
||||
|
||||
def shard_model(model,
|
||||
device_id,
|
||||
param_dtype=torch.bfloat16,
|
||||
reduce_dtype=torch.float32,
|
||||
buffer_dtype=torch.float32,
|
||||
fsdp_group = ['blocks'],
|
||||
sharding_strategy=ShardingStrategy.FULL_SHARD,
|
||||
sync_module_states=False):
|
||||
wrap_modules = []
|
||||
for module_name in fsdp_group:
|
||||
if hasattr(model, module_name):
|
||||
if isinstance(getattr(model, module_name), (list, tuple, nn.ModuleList)):
|
||||
wrap_modules.extend([m for m in getattr(model, module_name)])
|
||||
else:
|
||||
wrap_modules.extend([getattr(model, module_name)])
|
||||
else:
|
||||
warnings.warn("Can't find module {} in model".format(module_name))
|
||||
return FSDP(
|
||||
module=model,
|
||||
process_group=None,
|
||||
sharding_strategy=sharding_strategy,
|
||||
auto_wrap_policy=partial(
|
||||
# size_based_auto_wrap_policy, min_num_params=int(1e6),
|
||||
lambda_auto_wrap_policy,
|
||||
lambda_fn=lambda m: m in wrap_modules),
|
||||
mixed_precision=MixedPrecision(param_dtype=param_dtype,
|
||||
reduce_dtype=reduce_dtype,
|
||||
buffer_dtype=buffer_dtype),
|
||||
device_id=device_id,
|
||||
sync_module_states=sync_module_states)
|
||||
|
||||
|
||||
def get_module(instance, sub_module):
|
||||
sub_module_list = sub_module.split('.')
|
||||
for sub_mod in sub_module_list:
|
||||
if sub_mod == '':
|
||||
continue
|
||||
if hasattr(instance, sub_mod):
|
||||
instance = getattr(instance, sub_mod)
|
||||
else:
|
||||
return None
|
||||
return instance
|
||||
|
||||
|
||||
def set_module(instance, sub_module, value):
|
||||
sub_module_list = sub_module.split('.')
|
||||
instance_list = []
|
||||
for sub_mod in sub_module_list:
|
||||
if hasattr(instance, sub_mod):
|
||||
instance_list.append((instance, sub_mod))
|
||||
instance = getattr(instance, sub_mod)
|
||||
instance = value
|
||||
if len(instance_list) > 0:
|
||||
for parents_instance, sub_mod in instance_list[::-1]:
|
||||
setattr(parents_instance, sub_mod, instance)
|
||||
instance = parents_instance
|
||||
|
||||
|
||||
@SOLVERS.register_class()
|
||||
class LatentDiffusionSolver(BaseSolver):
|
||||
para_dict = {
|
||||
@@ -61,6 +125,24 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
'description':
|
||||
f'The shard strategy for fsdp, select from {list(sharding_strategy_map.keys())}',
|
||||
},
|
||||
'FSDP_REDUCE_DTYPE': {
|
||||
'value': 'float32',
|
||||
'description': 'The dtype of reduce in FSDP.'
|
||||
},
|
||||
'FSDP_BUFFER_DTYPE': {
|
||||
'value': 'float32',
|
||||
'description': 'The dtype of buffer in FSDP.'
|
||||
},
|
||||
'FSDP_SHARD_MODULES': {
|
||||
'value': ['model'],
|
||||
'description': 'The modules to be sharded in FSDP.'
|
||||
},
|
||||
'SAVE_MODULES': {
|
||||
'value':
|
||||
None,
|
||||
'description':
|
||||
'The modules to be saved, default is None to save all modules in checkpoint file.'
|
||||
},
|
||||
'IMAGE_LOG_STEP': {
|
||||
'value': 2000,
|
||||
'description': 'The interval for image log.',
|
||||
@@ -112,6 +194,13 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
else:
|
||||
self.logger.info('Use default backend.')
|
||||
self.model_shard = cfg.get('SHARDING_STRATEGY', 'full_shard')
|
||||
self.reduce_dtype = getattr(torch,
|
||||
cfg.get('FSDP_REDUCE_DTYPE', 'float32'))
|
||||
self.buffer_dtype = getattr(torch,
|
||||
cfg.get('FSDP_BUFFER_DTYPE', 'float32'))
|
||||
self.shard_modules = cfg.get('FSDP_SHARD_MODULES', ['model'])
|
||||
self.save_modules = cfg.get('SAVE_MODULES', ['model'])
|
||||
self.train_modules = cfg.get('TRAIN_MODULES', ['model'])
|
||||
self.image_log_step = cfg.get('IMAGE_LOG_STEP', 2000)
|
||||
self._image_out = defaultdict(list)
|
||||
self.load_model_only = cfg.get('LOAD_MODEL_ONLY', False)
|
||||
@@ -120,6 +209,7 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
self.sample_args = cfg.get('SAMPLE_ARGS', None)
|
||||
self.tuner_cfg = cfg.get('TUNER', None)
|
||||
self.freeze_cfg = cfg.get('FREEZE', None)
|
||||
self.log_train_num = cfg.get("LOG_TRAIN_NUM", -1)
|
||||
|
||||
def set_up(self):
|
||||
self.construct_data()
|
||||
@@ -178,7 +268,7 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
self.model = self.model.to(we.device_id)
|
||||
|
||||
def init_lr(self):
|
||||
rescale_lr = self.cfg.get('RESCALE_LR', True)
|
||||
rescale_lr = self.cfg.get('RESCALE_LR', False)
|
||||
if rescale_lr and 'train' in self.datas and self.cfg.have('OPTIMIZER'):
|
||||
if we.world_size > 1:
|
||||
all_batch_size = self.datas['train'].batch_size * we.world_size
|
||||
@@ -188,32 +278,76 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
self.cfg.OPTIMIZER.LEARNING_RATE /= 640
|
||||
|
||||
def init_opti(self):
|
||||
if hasattr(self.model, 'ignored_parameters'):
|
||||
train_params, ignored_params = self.model.parameters(
|
||||
), self.model.ignored_parameters()
|
||||
else:
|
||||
train_params, ignored_params = self.model.parameters(), None
|
||||
import torch.cuda.amp as amp
|
||||
|
||||
if we.is_distributed:
|
||||
if self.use_fairscale:
|
||||
from fairscale.nn.data_parallel import ShardedDataParallel
|
||||
from fairscale.optim.oss import OSS
|
||||
if hasattr(self.model, 'ignored_parameters'):
|
||||
train_params, ignored_params = self.model.parameters(
|
||||
), self.model.ignored_parameters()
|
||||
else:
|
||||
train_params, ignored_params = self.model.parameters(
|
||||
), None
|
||||
self.optimizer = OSS(params=train_params,
|
||||
optim=torch.optim.AdamW,
|
||||
lr=self.cfg.OPTIMIZER.LEARNING_RATE)
|
||||
self.model = ShardedDataParallel(self.model, self.optimizer)
|
||||
elif self.use_fsdp:
|
||||
mixed_precision = MixedPrecision(param_dtype=self.dtype,
|
||||
reduce_dtype=self.dtype,
|
||||
buffer_dtype=self.dtype)
|
||||
sharding_strategy = sharding_strategy_map[self.model_shard]
|
||||
self.model = FullyShardedDataParallel(
|
||||
self.model,
|
||||
mixed_precision=mixed_precision,
|
||||
cpu_offload=CPUOffload(offload_params=False),
|
||||
sharding_strategy=sharding_strategy,
|
||||
backward_prefetch=BackwardPrefetch.BACKWARD_PRE,
|
||||
device_id=torch.cuda.current_device(),
|
||||
ignored_parameters=ignored_params)
|
||||
shard_fn = partial
|
||||
if self.shard_modules is not None:
|
||||
for module in self.shard_modules:
|
||||
if isinstance(module, str):
|
||||
sub_module = get_module(self.model, module)
|
||||
if sub_module is not None:
|
||||
sub_module = shard_model(
|
||||
sub_module,
|
||||
device_id=we.device_id,
|
||||
param_dtype=self.dtype,
|
||||
reduce_dtype=self.reduce_dtype,
|
||||
buffer_dtype=self.buffer_dtype,
|
||||
sharding_strategy=sharding_strategy_map[self.model_shard],
|
||||
sync_module_states=True)
|
||||
set_module(self.model, module, sub_module)
|
||||
elif isinstance(module, (dict, Config)):
|
||||
sub_module = get_module(self.model, module["MODULE"])
|
||||
if sub_module is not None:
|
||||
sub_module = shard_model(
|
||||
sub_module,
|
||||
device_id=we.device_id,
|
||||
param_dtype=self.dtype,
|
||||
reduce_dtype=self.reduce_dtype,
|
||||
buffer_dtype=self.buffer_dtype,
|
||||
fsdp_group=module.get("FSDP_GROUP", ["blocks"]),
|
||||
sharding_strategy=sharding_strategy_map[self.model_shard],
|
||||
sync_module_states=True)
|
||||
set_module(self.model, module["MODULE"], sub_module)
|
||||
else:
|
||||
self.logger.warning(
|
||||
'FSDP_SHARD_MODULES is None, which means wraping the whold model as the '
|
||||
'fsdp instance. When using FSDP, it is necessary to specify the modules '
|
||||
'to be wrapped; otherwise, there may be a situation where submodules are '
|
||||
'not the root module, which can lead to unexpected issues. Specify the '
|
||||
'modules to be wrapped by setting FSDP_SHARD_MODULES to a list of modules '
|
||||
'that need wrapping.')
|
||||
self.model = shard_fn(self.model)
|
||||
train_params = []
|
||||
if self.train_modules is None:
|
||||
self.logger.warning(
|
||||
'When using FSDP, it is necessary to explicitly specify the modules to be '
|
||||
'trained or the modules for which gradients will be computed, otherwise, '
|
||||
'there will be issues with gradient calculation.')
|
||||
assert self.train_modules is None
|
||||
else:
|
||||
self.logger.info(
|
||||
f"The modules {','.join(self.train_modules)} 's parameters will be backwarded."
|
||||
)
|
||||
for module in self.train_modules:
|
||||
if hasattr(self.model, module):
|
||||
current_module = getattr(self.model, module)
|
||||
train_params += list(current_module.parameters())
|
||||
|
||||
self.optimizer = OPTIMIZERS.build(self.cfg.OPTIMIZER,
|
||||
logger=self.logger,
|
||||
parameters=train_params)
|
||||
@@ -223,14 +357,16 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
self.model,
|
||||
device_ids=[torch.cuda.current_device()],
|
||||
output_device=torch.cuda.current_device(),
|
||||
find_unused_parameters=True)
|
||||
self.optimizer = OPTIMIZERS.build(self.cfg.OPTIMIZER,
|
||||
logger=self.logger,
|
||||
parameters=train_params)
|
||||
find_unused_parameters=False)
|
||||
self.optimizer = OPTIMIZERS.build(
|
||||
self.cfg.OPTIMIZER,
|
||||
logger=self.logger,
|
||||
parameters=self.model.parameters())
|
||||
else:
|
||||
self.optimizer = OPTIMIZERS.build(self.cfg.OPTIMIZER,
|
||||
logger=self.logger,
|
||||
parameters=train_params)
|
||||
self.optimizer = OPTIMIZERS.build(
|
||||
self.cfg.OPTIMIZER,
|
||||
logger=self.logger,
|
||||
parameters=self.model.parameters())
|
||||
|
||||
if self.cfg.have('LR_SCHEDULER') and self.optimizer is not None:
|
||||
self.cfg.LR_SCHEDULER.TOTAL_STEPS = self.max_steps
|
||||
@@ -238,22 +374,22 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
logger=self.logger,
|
||||
optimizer=self.optimizer)
|
||||
|
||||
if self.cfg.DTYPE == 'float16':
|
||||
if self.cfg.DTYPE in ['float16']:
|
||||
if we.is_distributed:
|
||||
if self.use_fairscale:
|
||||
from fairscale.optim.grad_scaler import ShardedGradScaler
|
||||
self.scaler = ShardedGradScaler(enabled=True)
|
||||
elif self.use_fsdp:
|
||||
from torch.distributed.fsdp.sharded_grad_scaler import \
|
||||
ShardedGradScaler
|
||||
self.scaler = ShardedGradScaler()
|
||||
from torch.distributed.fsdp.sharded_grad_scaler import ShardedGradScaler
|
||||
self.scaler = ShardedGradScaler(enabled=True,
|
||||
process_group=None)
|
||||
else:
|
||||
self.scaler = amp.GradScaler()
|
||||
else:
|
||||
self.scaler = amp.GradScaler()
|
||||
else:
|
||||
self.scaler = None
|
||||
|
||||
self.logger.info(self.model)
|
||||
def load_checkpoint(self, checkpoint: dict):
|
||||
"""
|
||||
Load checkpoint function
|
||||
@@ -262,19 +398,56 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
"""
|
||||
if 'model' in checkpoint:
|
||||
if hasattr(self.model, 'module'):
|
||||
self.model.module.load_state_dict(checkpoint['model'])
|
||||
if self.save_modules is not None:
|
||||
for module in self.save_modules:
|
||||
current_module = get_module(self.model.module, module)
|
||||
if current_module is not None and module in checkpoint[
|
||||
'model']:
|
||||
current_module.load_state_dict(
|
||||
checkpoint['model'][module])
|
||||
self.logger.info(
|
||||
f'Load checkpoint for model.{module}')
|
||||
else:
|
||||
self.model.module.load_state_dict(checkpoint['model'])
|
||||
self.logger.info('Load checkpoint for model.')
|
||||
else:
|
||||
self.model.load_state_dict(checkpoint['model'])
|
||||
if self.save_modules is not None:
|
||||
for module in self.save_modules:
|
||||
current_module = get_module(self.model, module)
|
||||
self.logger.info(f'Load checkpoint for model.{module}')
|
||||
if current_module is not None and module in checkpoint[
|
||||
'model']:
|
||||
current_module.load_state_dict(
|
||||
checkpoint['model'][module])
|
||||
else:
|
||||
self.model.load_state_dict(checkpoint['model'])
|
||||
self.logger.info('Load checkpoint for model.')
|
||||
else:
|
||||
if hasattr(self.model, 'module'):
|
||||
self.model.module.load_state_dict(checkpoint)
|
||||
else:
|
||||
self.model.load_state_dict(checkpoint)
|
||||
self.logger.info('Load checkpoint for model.')
|
||||
if not self.load_model_only:
|
||||
if 'optimizer' in checkpoint and self.optimizer:
|
||||
self.optimizer.load_state_dict(checkpoint['optimizer'])
|
||||
if self.use_fsdp:
|
||||
for module in self.train_modules:
|
||||
current_module = get_module(self.model, module)
|
||||
if current_module is not None and module in checkpoint[
|
||||
'optimizer']:
|
||||
state = FSDP.optim_state_dict_to_load(
|
||||
current_module, self.optimizer,
|
||||
checkpoint['optimizer'][module])
|
||||
self.optimizer.load_state_dict(state)
|
||||
self.logger.info(
|
||||
f'Load checkpoint for optimizer {module}.')
|
||||
else:
|
||||
self.optimizer.load_state_dict(checkpoint['optimizer'])
|
||||
self.logger.info(f'Load checkpoint for optimizer.')
|
||||
if 'scaler' in checkpoint and self.scaler:
|
||||
self.scaler.load_state_dict(checkpoint['scaler'])
|
||||
self.logger.info(f'Load checkpoint for scaler.')
|
||||
self.logger.info('Load checkpoint finished.')
|
||||
|
||||
def save_checkpoint(self) -> dict:
|
||||
"""
|
||||
@@ -282,23 +455,76 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
:return:
|
||||
"""
|
||||
ckpt = dict()
|
||||
if not self.use_fsdp and not we.rank == 0:
|
||||
return ckpt
|
||||
ckpt['model'] = OrderedDict()
|
||||
if we.is_distributed:
|
||||
if self.use_fsdp:
|
||||
save_policy = FullStateDictConfig(offload_to_cpu=True,
|
||||
rank0_only=True)
|
||||
with FullyShardedDataParallel.state_dict_type(
|
||||
self.model, StateDictType.FULL_STATE_DICT,
|
||||
save_policy):
|
||||
ckpt['model'] = self.model.state_dict()
|
||||
if self.shard_modules is not None:
|
||||
if self.save_modules is None:
|
||||
self.logger.warning(
|
||||
'When using FSDP, after specifying the modules to be wrapped '
|
||||
'using FSDP_SHARD_MODULES, please set the modules to be saved '
|
||||
'in SAVE_MODULES. If set to None, it means all modules will '
|
||||
'be saved. However, for nested modules, the system cannot '
|
||||
'determine whether they are FSDP instances and need to be '
|
||||
'explicitly set.')
|
||||
assert self.save_modules is not None
|
||||
for module in self.save_modules:
|
||||
current_module = get_module(self.model, module)
|
||||
if current_module is not None:
|
||||
if isinstance(current_module, FSDP):
|
||||
# print(module, current_module._is_root)
|
||||
with FSDP.state_dict_type(
|
||||
current_module,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
save_policy):
|
||||
ckpt['model'][
|
||||
module] = current_module.state_dict()
|
||||
else:
|
||||
ckpt['model'][
|
||||
module] = current_module.state_dict()
|
||||
else:
|
||||
with FSDP.state_dict_type(self.model,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
save_policy):
|
||||
ckpt['model'] = self.model.state_dict()
|
||||
else:
|
||||
if hasattr(self.model, 'module'):
|
||||
ckpt['model'] = self.model.module.state_dict()
|
||||
model = self.model.module
|
||||
else:
|
||||
ckpt['model'] = self.model.state_dict()
|
||||
model = self.model
|
||||
if self.save_modules is not None:
|
||||
for module in self.save_modules:
|
||||
current_module = get_module(self.model, module)
|
||||
if current_module is not None:
|
||||
ckpt['model'][module] = current_module.state_dict()
|
||||
else:
|
||||
ckpt['model'] = model.state_dict()
|
||||
else:
|
||||
ckpt['model'] = self.model.state_dict()
|
||||
if hasattr(self.model, 'module'):
|
||||
model = self.model.module
|
||||
else:
|
||||
model = self.model
|
||||
if self.save_modules is not None:
|
||||
for module in self.save_modules:
|
||||
current_module = get_module(self.model, module)
|
||||
if current_module is not None:
|
||||
ckpt['model'][module] = current_module.state_dict()
|
||||
else:
|
||||
ckpt['model'] = model.state_dict()
|
||||
if self.optimizer and not self.use_fairscale:
|
||||
ckpt['optimizer'] = self.optimizer.state_dict()
|
||||
if self.use_fsdp and we.is_distributed:
|
||||
ckpt['optimizer'] = OrderedDict()
|
||||
for module in self.train_modules:
|
||||
if hasattr(self.model, module):
|
||||
current_module = getattr(self.model, module)
|
||||
ckpt['optimizer'][module] = FSDP.optim_state_dict(
|
||||
current_module, self.optimizer)
|
||||
else:
|
||||
ckpt['optimizer'] = self.optimizer.state_dict()
|
||||
if self.scaler:
|
||||
ckpt['scaler'] = self.scaler.state_dict()
|
||||
return ckpt
|
||||
@@ -354,7 +580,8 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
})
|
||||
self.current_batch_data[self.mode] = batch_data
|
||||
if self.sample_args:
|
||||
batch_data.update(self.sample_args.get_lowercase_dict())
|
||||
self.current_batch_data[self.mode].update(
|
||||
self.sample_args.get_lowercase_dict())
|
||||
with torch.autocast(device_type='cuda',
|
||||
enabled=self.use_amp,
|
||||
dtype=self.dtype):
|
||||
@@ -395,58 +622,29 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
log_data, log_label, ori_label = [], [], []
|
||||
for result in all_results:
|
||||
# the inference image use
|
||||
ret_images, ret_labels = [], []
|
||||
if 'hint' in result:
|
||||
merge_image = torch.cat([
|
||||
result['hint'][:result['image'].shape[0]], result['image']
|
||||
],
|
||||
dim=2)
|
||||
log_data.append((merge_image.permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
else:
|
||||
log_data.append(
|
||||
(result['image'].permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
log_label.append(result['prompt'] + ' NegPrompt: ' +
|
||||
result['n_prompt'])
|
||||
ret_images.append((result['hint'][:result['image'].shape[0]].permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
ret_labels.append(f"Control Image")
|
||||
ret_images.append(
|
||||
(result['image'].permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
ret_labels.append(result['prompt'] +
|
||||
" <font color='red'> |NegPrompt| </font> " +
|
||||
result['n_prompt'])
|
||||
log_data.append(ret_images)
|
||||
log_label.append(ret_labels)
|
||||
ori_label.append(result['prompt'])
|
||||
|
||||
self.register_probe({'test_label': log_label})
|
||||
self.register_probe({'eval_label': log_label})
|
||||
self.register_probe({
|
||||
'test_image':
|
||||
'eval_image':
|
||||
ProbeData(log_data,
|
||||
is_image=True,
|
||||
build_html=True,
|
||||
build_label=log_label)
|
||||
})
|
||||
|
||||
log_data, log_label, ori_label = [], [], []
|
||||
for result in all_results:
|
||||
# the inference image use
|
||||
if 'train_n_image' in result:
|
||||
if 'hint' in result:
|
||||
merge_image = torch.cat([
|
||||
result['hint'][:result['train_n_image'].shape[0]],
|
||||
result['train_n_image']
|
||||
],
|
||||
dim=2)
|
||||
log_data.append(
|
||||
(merge_image.permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
else:
|
||||
log_data.append((result['train_n_image'].permute(
|
||||
1, 2, 0).cpu().numpy() * 255).astype(np.uint8))
|
||||
log_label.append(result['prompt'] + 'NegPrompt' +
|
||||
result['train_n_prompt'])
|
||||
ori_label.append(result['prompt'])
|
||||
if len(log_data) > 0:
|
||||
self.register_probe({'test_train_n_label': log_label})
|
||||
self.register_probe({
|
||||
'test_train_n_image':
|
||||
ProbeData(log_data,
|
||||
is_image=True,
|
||||
build_html=True,
|
||||
build_label=log_label)
|
||||
})
|
||||
self.after_all_iter(self.hooks_dict[self._mode])
|
||||
|
||||
@torch.no_grad()
|
||||
@@ -468,14 +666,23 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
rank=we.rank)
|
||||
all_results.extend(results)
|
||||
self.after_iter(self.hooks_dict[self._mode])
|
||||
log_data, log_label = [], []
|
||||
log_data, log_label, ori_label = [], [], []
|
||||
for result in all_results:
|
||||
# the inference image use
|
||||
log_data.append((result['image'].permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
log_label.append(result['prompt'] +
|
||||
ret_images, ret_labels = [], []
|
||||
if 'hint' in result:
|
||||
ret_images.append((result['hint'][:result['image'].shape[0]].permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
ret_labels.append(f"Control Image")
|
||||
ret_images.append(
|
||||
(result['image'].permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
ret_labels.append(result['prompt'] +
|
||||
" <font color='red'> |NegPrompt| </font> " +
|
||||
result['n_prompt'])
|
||||
log_data.append(ret_images)
|
||||
log_label.append(ret_labels)
|
||||
ori_label.append(result['prompt'])
|
||||
|
||||
self.register_probe({'test_label': log_label})
|
||||
self.register_probe({
|
||||
@@ -485,27 +692,6 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
build_html=True,
|
||||
build_label=log_label)
|
||||
})
|
||||
|
||||
log_data, log_label = [], []
|
||||
for result in all_results:
|
||||
# the inference image use
|
||||
if 'train_n_image' in result:
|
||||
log_data.append(
|
||||
(result['train_n_image'].permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
log_label.append(result['prompt'] +
|
||||
" <font color='red'> |NegPrompt| </font> " +
|
||||
result['train_n_prompt'])
|
||||
if len(log_data) > 0:
|
||||
self.register_probe({'test_train_n_label': log_label})
|
||||
self.register_probe({
|
||||
'test_train_n_image':
|
||||
ProbeData(log_data,
|
||||
is_image=True,
|
||||
build_html=True,
|
||||
build_label=log_label)
|
||||
})
|
||||
|
||||
self.after_all_iter(self.hooks_dict[self._mode])
|
||||
|
||||
def add_tuner(self, tuner_cfg, model=None):
|
||||
@@ -524,7 +710,7 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
from swift import Swift
|
||||
model = Swift.prepare_model(self.model, config=swift_cfg_dict)
|
||||
|
||||
# self.logger.info([(key, param.shape) for key, param in self.model.named_parameters() if param.requires_grad])
|
||||
self.logger.info([(key, param.shape) for key, param in model.named_parameters() if param.requires_grad])
|
||||
return model
|
||||
|
||||
def freeze(self, freeze_cfg, model=None):
|
||||
@@ -582,6 +768,11 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
freeze_flag = sum([p in name for p in train_part]) > 0
|
||||
if freeze_flag:
|
||||
param.requires_grad = True
|
||||
elif isinstance(train_part, str):
|
||||
for name, param in freeze_model.named_parameters():
|
||||
if re.match(train_part, name):
|
||||
param.requires_grad = True
|
||||
self.logger.info([(key, param.shape) for key, param in freeze_model.named_parameters() if param.requires_grad])
|
||||
return model
|
||||
|
||||
@torch.no_grad()
|
||||
@@ -602,7 +793,11 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
return dict_to_yaml('solvername',
|
||||
__class__.__name__,
|
||||
LatentDiffusionSolver.para_dict,
|
||||
set_name=True)
|
||||
set_name=True,
|
||||
exclude_keys=[
|
||||
'EXTRA_KEYS', 'TRAIN_PRECISION', 'MAX_EPOCHS',
|
||||
'NUM_FOLDS'
|
||||
])
|
||||
|
||||
@property
|
||||
def image_out(self):
|
||||
@@ -620,29 +815,32 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
@property
|
||||
def probe_data(self):
|
||||
if not we.debug and self.mode == 'train':
|
||||
batch_data = transfer_data_to_cuda(self.current_batch_data[self.mode])
|
||||
self.eval_mode()
|
||||
with torch.autocast(device_type='cuda',
|
||||
enabled=self.use_amp,
|
||||
dtype=self.dtype):
|
||||
outputs = self.log_image(
|
||||
transfer_data_to_cuda(self.current_batch_data[self.mode]))
|
||||
batch_data['log_num'] = self.log_train_num
|
||||
results = self.run_step_eval(batch_data)
|
||||
images = batch_data['image'] if 'image' in batch_data else [None] * len(results)
|
||||
self.train_mode()
|
||||
log_data, log_label = [], []
|
||||
for result in outputs:
|
||||
for result, image in zip(results, images):
|
||||
ret_images, ret_labels = [], []
|
||||
if 'hint' in result:
|
||||
merge_image = torch.cat([
|
||||
result['orig'],
|
||||
result['hint'][:result['orig'].shape[0]],
|
||||
result['recon']
|
||||
],
|
||||
dim=2)
|
||||
else:
|
||||
merge_image = torch.cat([result['orig'], result['recon']],
|
||||
dim=2)
|
||||
log_data.append((merge_image.permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
log_label.append('recon image: ' + result['prompt'] +
|
||||
" <font color='red'> |NegPrompt| </font> " +
|
||||
result['n_prompt'])
|
||||
ret_images.append((result['hint'][:result['image'].shape[0]].permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
if image is not None:
|
||||
image = torch.clamp((image + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
ret_images.append((image.permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8))
|
||||
ret_labels.append(f'target image')
|
||||
|
||||
ret_images.append((result['image'].permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8))
|
||||
ret_labels.append(result['prompt']
|
||||
+ " <font color='red'> |NegPrompt| </font> "
|
||||
+ result['n_prompt'])
|
||||
log_data.append(ret_images)
|
||||
log_label.append(ret_labels)
|
||||
self.register_probe({
|
||||
'train_image':
|
||||
ProbeData(log_data,
|
||||
@@ -651,38 +849,6 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
build_label=log_label)
|
||||
})
|
||||
self.register_probe({'train_label': log_label})
|
||||
|
||||
# the inference image use
|
||||
log_data, log_label = [], []
|
||||
for result in outputs:
|
||||
if 'train_n_image' in result:
|
||||
if 'hint' in result:
|
||||
merge_image = torch.cat([
|
||||
result['orig'],
|
||||
result['hint'][:result['orig'].shape[0]],
|
||||
result['train_n_image']
|
||||
],
|
||||
dim=2)
|
||||
else:
|
||||
merge_image = torch.cat(
|
||||
[result['orig'], result['train_n_image']], dim=2)
|
||||
log_data.append(
|
||||
(merge_image.permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
log_label.append(
|
||||
'recon image: ' + result['prompt'] +
|
||||
" <font color='red'> |NegPrompt| </font> " +
|
||||
result['train_n_prompt'])
|
||||
|
||||
if len(log_data) > 0:
|
||||
self.register_probe({'train_n_label': log_label})
|
||||
self.register_probe({
|
||||
'train_n_image':
|
||||
ProbeData(log_data,
|
||||
is_image=True,
|
||||
build_html=True,
|
||||
build_label=log_label)
|
||||
})
|
||||
return super().probe_data
|
||||
|
||||
def print_memory_status(self):
|
||||
@@ -719,12 +885,20 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
if logger is None:
|
||||
logger = self.logger
|
||||
train_param_dict = {}
|
||||
forzen_param_dict = {}
|
||||
frozen_param_dict = {}
|
||||
ema_param_dict = {}
|
||||
all_param_numel = 0
|
||||
if we.debug:
|
||||
for key, _ in model.named_modules():
|
||||
logger.info(f'sub modules {key}.')
|
||||
for key, val in model.named_parameters():
|
||||
if 'ema' in key:
|
||||
sub_key = '.'.join(key.split('.', 1)[-1].split('.', 1)[:1])
|
||||
if sub_key in ema_param_dict:
|
||||
ema_param_dict[sub_key] += val.numel()
|
||||
else:
|
||||
ema_param_dict[sub_key] = val.numel()
|
||||
continue
|
||||
if val.requires_grad:
|
||||
sub_key = '.'.join(key.split('.', 1)[-1].split('.', 2)[:2])
|
||||
if sub_key in train_param_dict:
|
||||
@@ -733,20 +907,26 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
train_param_dict[sub_key] = val.numel()
|
||||
else:
|
||||
sub_key = '.'.join(key.split('.', 1)[-1].split('.', 1)[:1])
|
||||
if sub_key in forzen_param_dict:
|
||||
forzen_param_dict[sub_key] += val.numel()
|
||||
if sub_key in frozen_param_dict:
|
||||
frozen_param_dict[sub_key] += val.numel()
|
||||
else:
|
||||
forzen_param_dict[sub_key] = val.numel()
|
||||
frozen_param_dict[sub_key] = val.numel()
|
||||
all_param_numel += val.numel()
|
||||
if we.debug:
|
||||
logger.info(key)
|
||||
train_param_numel = sum(train_param_dict.values())
|
||||
forzen_param_numel = sum(forzen_param_dict.values())
|
||||
frozen_param_numel = sum(frozen_param_dict.values())
|
||||
logger.info(
|
||||
f'Load trainable params {train_param_numel} / {all_param_numel} = '
|
||||
f'{train_param_numel / all_param_numel:.2%}, '
|
||||
f'train part: {train_param_dict}.')
|
||||
logger.info(
|
||||
f'Load forzen params {forzen_param_numel} / {all_param_numel} = '
|
||||
f'{forzen_param_numel / all_param_numel:.2%}, '
|
||||
f'forzen part: {forzen_param_dict}.')
|
||||
f'Load frozen params {frozen_param_numel} / {all_param_numel} = '
|
||||
f'{frozen_param_numel / all_param_numel:.2%}, '
|
||||
f'frozen part: {frozen_param_dict}.')
|
||||
if len(ema_param_dict) > 0:
|
||||
ema_param_numel = sum(ema_param_dict.values())
|
||||
logger.info(
|
||||
f'Load ema frozen params {ema_param_numel} / {all_param_numel} = '
|
||||
f'{ema_param_numel / all_param_numel:.2%}, '
|
||||
f'frozen part: {ema_param_dict}.')
|
||||
@@ -1,8 +1,13 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import os
|
||||
import warnings
|
||||
|
||||
import torch
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
from scepter.modules.utils.distribute import we
|
||||
|
||||
from scepter.modules.solver.hooks.hook import Hook
|
||||
from scepter.modules.solver.hooks.registry import HOOKS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
@@ -30,6 +35,30 @@ class BackwardHook(Hook):
|
||||
'EMPTY_CACHE_STEP': {
|
||||
'value': -1,
|
||||
'description': 'the memory empty step!'
|
||||
},
|
||||
'DO_PROFILE': {
|
||||
'value': False,
|
||||
'description': 'whether to do profiling!'
|
||||
},
|
||||
'PROFILE_DIR': {
|
||||
'value': None,
|
||||
'description': 'the dir for profiling!'
|
||||
},
|
||||
'PROFILE_WAIT': {
|
||||
'value': 1,
|
||||
'description': 'the wait steps for profiling!'
|
||||
},
|
||||
'PROFILE_WARMUP': {
|
||||
'value': 1,
|
||||
'description': 'the warmup steps for profiling!'
|
||||
},
|
||||
'PROFILE_ACTIVE': {
|
||||
'value': 3,
|
||||
'description': 'the active steps for profiling!'
|
||||
},
|
||||
'REPEAT': {
|
||||
'value': 1,
|
||||
'description': 'the repeat steps for profiling!'
|
||||
}
|
||||
}]
|
||||
|
||||
@@ -40,17 +69,48 @@ class BackwardHook(Hook):
|
||||
self.empty_cache_step = cfg.get('EMPTY_CACHE_STEP', -1)
|
||||
self.accumulate_step = cfg.get('ACCUMULATE_STEP', 1)
|
||||
self.current_step = 0
|
||||
self.wait = cfg.get('PROFILE_WAIT', 1)
|
||||
self.warmup = cfg.get('PROFILE_WARMUP', 1)
|
||||
self.active = cfg.get('PROFILE_ACTIVE', 3)
|
||||
self.repeat = cfg.get('REPEAT', 1)
|
||||
self.do_profile = cfg.get('DO_PROFILE', False)
|
||||
self.profile_dir = cfg.get('PROFILE_DIR', None)
|
||||
self.profile_step = 0
|
||||
self.prof = None
|
||||
|
||||
def before_solve(self, solver):
|
||||
if we.rank != 0:
|
||||
return
|
||||
if self.profile_dir is None:
|
||||
self.log_dir = os.path.join(solver.work_dir, 'profile')
|
||||
self._local_log_dir, _ = FS.map_to_local(self.log_dir)
|
||||
os.makedirs(self._local_log_dir, exist_ok=True)
|
||||
if self.do_profile:
|
||||
self.prof = torch.profiler.profile(
|
||||
schedule=torch.profiler.schedule(wait=self.wait, warmup=self.warmup, active=self.active, repeat=self.repeat),
|
||||
on_trace_ready=torch.profiler.tensorboard_trace_handler(self._local_log_dir),
|
||||
record_shapes=True,
|
||||
with_stack=True)
|
||||
self.prof.start()
|
||||
solver.logger.info(f'Profiler start ...')
|
||||
solver.logger.info(f'Profiler: save to {self.log_dir}')
|
||||
def profile(self, solver):
|
||||
if self.prof is None: return
|
||||
if we.rank == 0 and self.do_profile:
|
||||
if self.profile_step < self.wait + self.warmup + self.active:
|
||||
self.prof.step()
|
||||
self.profile_step += 1
|
||||
else:
|
||||
self.prof.stop()
|
||||
self.do_profile = False
|
||||
solver.logger.info(f'Profiler stop after {self.profile_step} steps')
|
||||
FS.put_dir_from_local_dir(self._local_log_dir, self.log_dir)
|
||||
def grad_clip(self, parameters):
|
||||
torch.nn.utils.clip_grad_norm_(parameters=parameters,
|
||||
max_norm=self.gradient_clip,
|
||||
norm_type=2)
|
||||
|
||||
def after_iter(self, solver):
|
||||
if (hasattr(solver, 'use_fsdp')
|
||||
and solver.use_fsdp) and self.accumulate_step > 1:
|
||||
self.logger.info("Fsdp don't surpport gradient accumulate.")
|
||||
self.accumulate_step = 1
|
||||
if solver.optimizer is not None and solver.is_train_mode:
|
||||
if solver.loss is None:
|
||||
warnings.warn(
|
||||
@@ -58,28 +118,37 @@ class BackwardHook(Hook):
|
||||
)
|
||||
return
|
||||
if solver.scaler is not None:
|
||||
solver.scaler.scale(solver.loss).backward()
|
||||
solver.scaler.scale(solver.loss/self.accumulate_step).backward()
|
||||
if self.gradient_clip > 0:
|
||||
solver.scaler.unscale_(solver.optimizer)
|
||||
self.grad_clip(solver.train_parameters())
|
||||
self.current_step += 1
|
||||
# Suppose profiler run after backward, so we need to set backward_prev_step
|
||||
# as the previous one step before the backward step
|
||||
if self.current_step % self.accumulate_step == 0:
|
||||
self.profile(solver)
|
||||
solver.scaler.step(solver.optimizer)
|
||||
solver.scaler.update()
|
||||
solver.optimizer.zero_grad()
|
||||
else:
|
||||
solver.loss.backward()
|
||||
(solver.loss/self.accumulate_step).backward()
|
||||
if self.gradient_clip > 0:
|
||||
self.grad_clip(solver.train_parameters())
|
||||
self.current_step += 1
|
||||
# Suppose profiler run after backward, so we need to set backward_prev_step
|
||||
# as the previous one step before the backward step
|
||||
if self.current_step % self.accumulate_step == 0:
|
||||
self.profile(solver)
|
||||
solver.optimizer.step()
|
||||
solver.optimizer.zero_grad()
|
||||
if solver.lr_scheduler:
|
||||
if self.current_step % self.accumulate_step == 0:
|
||||
solver.lr_scheduler.step()
|
||||
if self.current_step % self.accumulate_step == 0:
|
||||
setattr(solver, 'backward_step', True)
|
||||
self.current_step = 0
|
||||
else:
|
||||
setattr(solver, 'backward_step', False)
|
||||
solver.loss = None
|
||||
if self.empty_cache_step > 0 and solver.total_iter % self.empty_cache_step == 0:
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
@@ -13,7 +13,6 @@ from scepter.modules.solver.hooks.registry import HOOKS
|
||||
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 swift import push_to_hub
|
||||
|
||||
_DEFAULT_CHECKPOINT_PRIORITY = 300
|
||||
|
||||
@@ -109,35 +108,42 @@ class CheckpointHook(Hook):
|
||||
or solver.total_iter == solver.max_steps - 1):
|
||||
solver.logger.info(
|
||||
f'Saving checkpoint after {solver.total_iter + 1} steps')
|
||||
if we.rank == 0:
|
||||
save_path = osp.join(
|
||||
solver.work_dir,
|
||||
'checkpoints/{}-{}.pth'.format(self.save_name_prefix,
|
||||
solver.total_iter + 1))
|
||||
if not self.disable_save_snapshot:
|
||||
save_path = osp.join(
|
||||
solver.work_dir,
|
||||
'checkpoints/{}-{}.pth'.format(self.save_name_prefix,
|
||||
solver.total_iter + 1))
|
||||
if not self.disable_save_snapshot:
|
||||
checkpoint = solver.save_checkpoint()
|
||||
if we.rank == 0:
|
||||
with FS.put_to(save_path) as local_path:
|
||||
with open(local_path, 'wb') as f:
|
||||
checkpoint = solver.save_checkpoint()
|
||||
torch.save(checkpoint, f)
|
||||
del checkpoint
|
||||
|
||||
from swift import SwiftModel
|
||||
if isinstance(solver.model, SwiftModel):
|
||||
save_path = osp.join(
|
||||
solver.work_dir,
|
||||
'checkpoints/{}-{}'.format(self.save_name_prefix,
|
||||
solver.total_iter + 1))
|
||||
from swift import SwiftModel
|
||||
if isinstance(solver.model, SwiftModel) or (
|
||||
hasattr(solver.model, 'module')
|
||||
and isinstance(solver.model.module, SwiftModel)):
|
||||
save_path = osp.join(
|
||||
solver.work_dir,
|
||||
'checkpoints/{}-{}'.format(self.save_name_prefix,
|
||||
solver.total_iter + 1))
|
||||
if we.rank == 0:
|
||||
local_folder, _ = FS.map_to_local(save_path)
|
||||
solver.model.save_pretrained(local_folder)
|
||||
if hasattr(solver.model, 'module'):
|
||||
solver.model.module.save_pretrained(local_folder)
|
||||
else:
|
||||
solver.model.save_pretrained(local_folder)
|
||||
FS.put_dir_from_local_dir(local_folder, save_path)
|
||||
else:
|
||||
if hasattr(solver, 'save_pretrained'):
|
||||
save_path = osp.join(
|
||||
solver.work_dir,
|
||||
'checkpoints/{}-{}'.format(self.save_name_prefix,
|
||||
solver.total_iter + 1))
|
||||
local_folder, _ = FS.map_to_local(save_path)
|
||||
FS.make_dir(local_folder)
|
||||
ckpt, cfg = solver.save_pretrained()
|
||||
else:
|
||||
if hasattr(solver, 'save_pretrained'):
|
||||
save_path = osp.join(
|
||||
solver.work_dir, 'checkpoints/{}-{}'.format(
|
||||
self.save_name_prefix, solver.total_iter + 1))
|
||||
local_folder, _ = FS.map_to_local(save_path)
|
||||
FS.make_dir(local_folder)
|
||||
ckpt, cfg = solver.save_pretrained()
|
||||
if we.rank == 0:
|
||||
with FS.put_to(
|
||||
os.path.join(
|
||||
local_folder,
|
||||
@@ -150,6 +156,7 @@ class CheckpointHook(Hook):
|
||||
'configuration.json')) as local_path:
|
||||
json.dump(cfg, open(local_path, 'w'))
|
||||
FS.put_dir_from_local_dir(local_folder, save_path)
|
||||
del ckpt
|
||||
|
||||
if self.save_last and solver.total_iter == solver.max_steps - 1:
|
||||
with FS.get_fs_client(save_path) as client:
|
||||
@@ -219,8 +226,10 @@ class CheckpointHook(Hook):
|
||||
with open(local_file, 'wb') as f:
|
||||
torch.save(checkpoint['pre_state_dict'], f)
|
||||
client.put_object_from_local_file(local_file, save_path)
|
||||
del checkpoint
|
||||
|
||||
def after_all_iter(self, solver):
|
||||
from swift import push_to_hub
|
||||
if we.rank == 0:
|
||||
if self.push_to_hub and self.last_ckpt:
|
||||
with FS.get_dir_to_local_dir(self.last_ckpt) as local_dir:
|
||||
|
||||
@@ -23,6 +23,26 @@ class ProbeDataHook(Hook):
|
||||
'PROB_INTERVAL': {
|
||||
'value': 1000,
|
||||
'description': 'the interval for log print!'
|
||||
},
|
||||
'SAVE_NAME_PREFIX': {
|
||||
'value': 'step',
|
||||
'description': 'the prefix for save name!'
|
||||
},
|
||||
'SAVE_PROBE_PREFIX': {
|
||||
'value': None,
|
||||
'description': 'the prefix for save probe!'
|
||||
},
|
||||
'SAVE_LAST': {
|
||||
'value': False,
|
||||
'description': 'whether to save last!'
|
||||
},
|
||||
'SAVE_IMAGE_POSTFIX': {
|
||||
'value': 'jpg',
|
||||
'description': 'the postfix for save image!'
|
||||
},
|
||||
'SAVE_VIDEO_POSTFIX': {
|
||||
'value': 'mp4',
|
||||
'description': 'the postfix for save video!'
|
||||
}
|
||||
}]
|
||||
|
||||
@@ -34,31 +54,46 @@ class ProbeDataHook(Hook):
|
||||
self.save_probe_prefix = cfg.get('SAVE_PROBE_PREFIX', None)
|
||||
self.save_last = cfg.get('SAVE_LAST', False)
|
||||
self.save_image_postfix = cfg.get('SAVE_IMAGE_POSTFIX', 'jpg')
|
||||
self.save_video_postfix = cfg.get('SAVE_VIDEO_POSTFIX', 'mp4')
|
||||
|
||||
def before_all_iter(self, solver):
|
||||
pass
|
||||
|
||||
if not solver.mode == 'train' and hasattr(solver, 'eval_interval'):
|
||||
solver.eval_interval = self.prob_interval
|
||||
def before_iter(self, solver):
|
||||
pass
|
||||
|
||||
def get_key_level_prefix(self, key, save_folder, total_iter):
|
||||
if self.save_probe_prefix is not None:
|
||||
ret_prefix = os.path.join(save_folder,
|
||||
self.save_probe_prefix)
|
||||
else:
|
||||
ret_prefix = os.path.join(
|
||||
save_folder,
|
||||
key.replace('/', '_') + f'_step_{total_iter}')
|
||||
return ret_prefix
|
||||
def after_iter(self, solver):
|
||||
if solver.mode == 'train' and solver.total_iter % self.prob_interval == 0:
|
||||
probe_dict = solver.probe_data
|
||||
if we.rank == 0:
|
||||
save_folder = os.path.join(
|
||||
solver.work_dir,
|
||||
f'{solver.mode}_probe/{self.save_name_prefix}-{solver.total_iter}'
|
||||
save_folder = os.path.join(
|
||||
solver.work_dir,
|
||||
f'{solver.mode}_probe/{self.save_name_prefix}-{solver.total_iter}'
|
||||
)
|
||||
setattr(solver,
|
||||
f'{solver.mode}_pre_save_paras',
|
||||
{
|
||||
"save_folder":save_folder,
|
||||
"save_probe_prefix": self.save_probe_prefix,
|
||||
"image_postfix": self.save_image_postfix,
|
||||
"video_postfix": self.save_video_postfix,
|
||||
"step": solver.total_iter,
|
||||
}
|
||||
)
|
||||
gather_probe_dict = solver.probe_data
|
||||
if we.rank == 0:
|
||||
ret_data = {}
|
||||
for k, v in probe_dict.items():
|
||||
if self.save_probe_prefix is not None:
|
||||
ret_prefix = os.path.join(save_folder,
|
||||
self.save_probe_prefix)
|
||||
else:
|
||||
ret_prefix = os.path.join(
|
||||
save_folder,
|
||||
k.replace('/', '_') + f'_step_{solver.total_iter}')
|
||||
ret_one = v.to_log(ret_prefix, self.save_image_postfix)
|
||||
for k, v in gather_probe_dict.items():
|
||||
ret_prefix = self.get_key_level_prefix(k, save_folder, solver.total_iter)
|
||||
ret_one = v.to_log(ret_prefix, image_postfix=self.save_image_postfix,
|
||||
video_postfix=self.save_video_postfix)
|
||||
if (isinstance(ret_one, list)
|
||||
or isinstance(ret_one, dict)) and len(ret_one) < 1:
|
||||
continue
|
||||
@@ -74,23 +109,29 @@ class ProbeDataHook(Hook):
|
||||
|
||||
def after_all_iter(self, solver):
|
||||
if not solver.mode == 'train':
|
||||
step = solver._total_iter[
|
||||
'train'] if 'train' in solver._total_iter else 0
|
||||
save_folder = os.path.join(
|
||||
solver.work_dir,
|
||||
f'{solver.mode}_probe/{self.save_name_prefix}-{step}'
|
||||
)
|
||||
setattr(solver,
|
||||
f'{solver.mode}_pre_save_paras',
|
||||
{
|
||||
"save_folder": save_folder,
|
||||
"save_probe_prefix": self.save_probe_prefix,
|
||||
"image_postfix": self.save_image_postfix,
|
||||
"step": step,
|
||||
"video_postfix": self.save_video_postfix
|
||||
}
|
||||
)
|
||||
probe_dict = solver.probe_data
|
||||
if we.rank == 0:
|
||||
step = solver._total_iter[
|
||||
'train'] if 'train' in solver._total_iter else 0
|
||||
save_folder = os.path.join(
|
||||
solver.work_dir,
|
||||
f'{solver.mode}_probe/{self.save_name_prefix}-{step}')
|
||||
ret_data = {}
|
||||
for k, v in probe_dict.items():
|
||||
if self.save_probe_prefix is not None:
|
||||
ret_prefix = os.path.join(save_folder,
|
||||
self.save_probe_prefix)
|
||||
else:
|
||||
ret_prefix = os.path.join(
|
||||
save_folder,
|
||||
k.replace('/', '_') + f'_step_{step}')
|
||||
ret_one = v.to_log(ret_prefix, self.save_image_postfix)
|
||||
ret_prefix = self.get_key_level_prefix(k, save_folder, step)
|
||||
ret_one = v.to_log(ret_prefix, image_postfix=self.save_image_postfix,
|
||||
video_postfix=self.save_video_postfix)
|
||||
if (isinstance(ret_one, list)
|
||||
or isinstance(ret_one, dict)) and len(ret_one) < 1:
|
||||
continue
|
||||
|
||||
@@ -124,12 +124,17 @@ class LogHook(Hook):
|
||||
|
||||
self.time = time.time()
|
||||
self.start_time = time.time()
|
||||
self.all_throughput = 0
|
||||
self.data_time = 0
|
||||
self.batch_size = defaultdict(dict)
|
||||
|
||||
def before_all_iter(self, solver):
|
||||
self.time = time.time()
|
||||
self.last_log_step = (solver.mode, 0)
|
||||
|
||||
if hasattr(solver, "datas"):
|
||||
for k, v in solver.datas.items():
|
||||
if hasattr(v, 'batch_size'):
|
||||
self.batch_size[k] = v.batch_size
|
||||
def before_iter(self, solver):
|
||||
data_time = time.time() - self.time
|
||||
self.data_time = data_time
|
||||
@@ -141,16 +146,21 @@ class LogHook(Hook):
|
||||
outputs = solver.iter_outputs.copy()
|
||||
outputs['time'] = iter_time
|
||||
outputs['data_time'] = self.data_time
|
||||
if 'batch_size' in outputs:
|
||||
batch_size = outputs.pop('batch_size')
|
||||
else:
|
||||
batch_size = 1
|
||||
if solver.mode in self.batch_size:
|
||||
outputs['throughput'] = int(self.batch_size[solver.mode] * we.world_size / iter_time * 86400)
|
||||
log_agg.update(outputs, 1)
|
||||
log_agg = log_agg.aggregate(self.log_interval)
|
||||
if 'throughput' in log_agg:
|
||||
log_agg['throughput'] = f"{int(log_agg['throughput'][-1])}/day"
|
||||
if solver.mode in self.batch_size:
|
||||
log_agg['all_throughput'] = (solver.iter + 1) * we.world_size * self.batch_size[solver.mode]
|
||||
|
||||
if self.show_gpu_mem:
|
||||
outputs['nvidia-smi'] = print_memory_status()
|
||||
log_agg.update(outputs, batch_size)
|
||||
log_agg['nvidia-smi'] = str(print_memory_status()) +"MiB"
|
||||
|
||||
if (solver.iter + 1) % self.log_interval == 0:
|
||||
_print_iter_log(solver,
|
||||
log_agg.aggregate(self.log_interval),
|
||||
log_agg,
|
||||
start_time=self.start_time,
|
||||
mode=solver.mode)
|
||||
self.last_log_step = (solver.mode, solver.iter + 1)
|
||||
|
||||
@@ -20,7 +20,7 @@ _SECURE_KEYWORDS = [
|
||||
_SECURE_VALUEWORDS = ['oss://', 'oss-'] # -> "#####"
|
||||
|
||||
|
||||
def dict_to_yaml(module_name, name, json_config, set_name=False):
|
||||
def dict_to_yaml(module_name, name, json_config, set_name=False, exclude_keys=[]):
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
@@ -70,6 +70,9 @@ def dict_to_yaml(module_name, name, json_config, set_name=False):
|
||||
yaml_str = ''
|
||||
# print(level_num, json_config)
|
||||
if isinstance(json_config, dict):
|
||||
for key in exclude_keys:
|
||||
if key in json_config:
|
||||
json_config.pop(key)
|
||||
if 'value' in json_config:
|
||||
value = json_config['value']
|
||||
if isinstance(value, dict):
|
||||
@@ -322,39 +325,39 @@ class Config(object):
|
||||
'CUDNN_DETERMINISTIC': True,
|
||||
'CUDNN_BENCHMARK': False
|
||||
}
|
||||
self.logger.info(
|
||||
f"ENV is not set and will use default ENV as {self.cfg_dict['ENV']}; "
|
||||
f'If want to change this value, please set them in your config.'
|
||||
)
|
||||
# self.logger.info(
|
||||
# f"ENV is not set and will use default ENV as {self.cfg_dict['ENV']}; "
|
||||
# f'If want to change this value, please set them in your config.'
|
||||
# )
|
||||
else:
|
||||
if 'SEED' not in self.cfg_dict['ENV']:
|
||||
self.cfg_dict['ENV']['SEED'] = 2023
|
||||
self.logger.info(
|
||||
f"SEED is not set and will use default SEED as {self.cfg_dict['ENV']['SEED']}; "
|
||||
f'If want to change this value, please set it in your config.'
|
||||
)
|
||||
# self.logger.info(
|
||||
# f"SEED is not set and will use default SEED as {self.cfg_dict['ENV']['SEED']}; "
|
||||
# f'If want to change this value, please set it in your config.'
|
||||
# )
|
||||
os.environ['ES_SEED'] = str(self.cfg_dict['ENV']['SEED'])
|
||||
self._update_dict(self.cfg_dict)
|
||||
if load:
|
||||
self.logger.info(f'Parse cfg file as \n {self.dump()}')
|
||||
# if load:
|
||||
# self.logger.info(f'Parse cfg file as \n {self.dump()}')
|
||||
|
||||
def load_from_file(self, file_name):
|
||||
self.logger.info(f'Loading config from {file_name}')
|
||||
# self.logger.info(f'Loading config from {file_name}')
|
||||
if file_name is None or not os.path.exists(file_name):
|
||||
self.logger.info(f'File {file_name} does not exist!')
|
||||
self.logger.warning(
|
||||
f"Cfg file is None or doesn't exist, Skip loading config from {file_name}."
|
||||
f"Cfg file is None or doesn't exist, Skip loading config from [{file_name}]."
|
||||
)
|
||||
return
|
||||
if file_name.endswith('.json'):
|
||||
self.cfg_dict = self._load_json(file_name)
|
||||
self.logger.info(
|
||||
f'System take {file_name} as json, because we find json in this file'
|
||||
f'Loading config from [{file_name}] as json file.'
|
||||
)
|
||||
elif file_name.endswith('.yaml'):
|
||||
self.cfg_dict = self._load_yaml(file_name)
|
||||
self.logger.info(
|
||||
f'System take {file_name} as yaml, because we find yaml in this file'
|
||||
f'Loading config from [{file_name}] as yaml file.'
|
||||
)
|
||||
else:
|
||||
self.logger.info(
|
||||
@@ -616,6 +619,20 @@ class Config(object):
|
||||
config_new[key] = val
|
||||
return config_new
|
||||
|
||||
def get_uppercase_dict(self, cfg_dict=None):
|
||||
if cfg_dict is None:
|
||||
cfg_dict = self.get_dict()
|
||||
config_new = {}
|
||||
for key, val in cfg_dict.items():
|
||||
if isinstance(key, str):
|
||||
if isinstance(val, dict):
|
||||
config_new[key.upper()] = self.get_uppercase_dict(val)
|
||||
else:
|
||||
config_new[key.upper()] = val
|
||||
else:
|
||||
config_new[key] = val
|
||||
return config_new
|
||||
|
||||
@staticmethod
|
||||
def get_plain_cfg(cfg=None):
|
||||
if isinstance(cfg, Config):
|
||||
|
||||
@@ -83,7 +83,13 @@ def transfer_data_to_cuda(data_map: dict) -> dict:
|
||||
elif isinstance(value, dict):
|
||||
ret[key] = transfer_data_to_cuda(value)
|
||||
elif isinstance(value, (list, tuple)):
|
||||
ret[key] = type(value)([transfer_data_to_cuda(t) for t in value])
|
||||
ret_data = []
|
||||
for t in value:
|
||||
if not isinstance(t, dict):
|
||||
ret_data.append(transfer_data_to_cuda({'data': t})['data'])
|
||||
else:
|
||||
ret_data.append(transfer_data_to_cuda(t))
|
||||
ret[key] = type(value)(ret_data)
|
||||
else:
|
||||
ret[key] = value
|
||||
return ret
|
||||
|
||||
@@ -4,12 +4,16 @@ import functools
|
||||
import os
|
||||
import pickle
|
||||
import random
|
||||
import socket
|
||||
import warnings
|
||||
from collections import OrderedDict
|
||||
from datetime import timedelta
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.autograd import Function
|
||||
|
||||
from scepter.modules.utils.model import StdMsg
|
||||
|
||||
__all__ = [
|
||||
@@ -202,6 +206,96 @@ def barrier():
|
||||
dist.barrier()
|
||||
|
||||
|
||||
def all_gather(tensor, uniform_size=True, group=None, **kwargs):
|
||||
world_size = dist.get_world_size(group)
|
||||
if world_size == 1:
|
||||
return [tensor]
|
||||
assert tensor.is_contiguous(), \
|
||||
'ops.all_gather requires the tensor to be contiguous()'
|
||||
|
||||
if uniform_size:
|
||||
tensor_list = [torch.empty_like(tensor) for _ in range(world_size)]
|
||||
dist.all_gather(tensor_list, tensor, group, **kwargs)
|
||||
return tensor_list
|
||||
else:
|
||||
# collect tensor shapes across GPUs
|
||||
shape = tuple(tensor.shape)
|
||||
shape_list = generalized_all_gather(shape, group)
|
||||
|
||||
# flatten the tensor
|
||||
tensor = tensor.reshape(-1)
|
||||
size = int(np.prod(shape))
|
||||
size_list = [int(np.prod(u)) for u in shape_list]
|
||||
max_size = max(size_list)
|
||||
|
||||
# pad to maximum size
|
||||
if size != max_size:
|
||||
padding = tensor.new_zeros(max_size - size)
|
||||
tensor = torch.cat([tensor, padding], dim=0)
|
||||
|
||||
# all_gather
|
||||
tensor_list = [torch.empty_like(tensor) for _ in range(world_size)]
|
||||
dist.all_gather(tensor_list, tensor, group, **kwargs)
|
||||
|
||||
# reshape tensors
|
||||
tensor_list = [
|
||||
t[:n].view(s)
|
||||
for t, n, s in zip(tensor_list, size_list, shape_list)
|
||||
]
|
||||
return tensor_list
|
||||
|
||||
|
||||
def _pad_to_largest_tensor(tensor, group):
|
||||
world_size = dist.get_world_size(group=group)
|
||||
assert world_size >= 1, \
|
||||
'gather/all_gather must be called from ranks within' \
|
||||
'the give group!'
|
||||
local_size = torch.tensor([tensor.numel()],
|
||||
dtype=torch.int64,
|
||||
device=tensor.device)
|
||||
size_list = [
|
||||
torch.zeros([1], dtype=torch.int64, device=tensor.device)
|
||||
for _ in range(world_size)
|
||||
]
|
||||
|
||||
# gather tensors and compute the maximum size
|
||||
dist.all_gather(size_list, local_size, group=group)
|
||||
size_list = [int(size.item()) for size in size_list]
|
||||
max_size = max(size_list)
|
||||
|
||||
# pad tensors to the same size
|
||||
if local_size != max_size:
|
||||
padding = torch.zeros((max_size - local_size, ),
|
||||
dtype=torch.uint8,
|
||||
device=tensor.device)
|
||||
tensor = torch.cat((tensor, padding), dim=0)
|
||||
return size_list, tensor
|
||||
|
||||
|
||||
def generalized_all_gather(data, group=None):
|
||||
if dist.get_world_size(group) == 1:
|
||||
return [data]
|
||||
if group is None:
|
||||
group = get_global_gloo_group()
|
||||
|
||||
tensor = _serialize_to_tensor(data, group)
|
||||
size_list, tensor = _pad_to_largest_tensor(tensor, group)
|
||||
max_size = max(size_list)
|
||||
|
||||
# receiving tensors from all ranks
|
||||
tensor_list = [
|
||||
torch.empty((max_size, ), dtype=torch.uint8, device=tensor.device)
|
||||
for _ in size_list
|
||||
]
|
||||
dist.all_gather(tensor_list, tensor, group=group)
|
||||
|
||||
data_list = []
|
||||
for size, tensor in zip(size_list, tensor_list):
|
||||
buffer = tensor.cpu().numpy().tobytes()[:size]
|
||||
data_list.append(pickle.loads(buffer))
|
||||
return data_list
|
||||
|
||||
|
||||
@functools.lru_cache()
|
||||
def get_global_gloo_group():
|
||||
backend = dist.get_backend()
|
||||
@@ -223,12 +317,14 @@ def reduce_scatter(output,
|
||||
|
||||
def all_reduce(tensor, op=dist.ReduceOp.SUM, group=None, **kwargs):
|
||||
if we.is_distributed:
|
||||
return dist.all_reduce(tensor, op, group, **kwargs)
|
||||
dist.all_reduce(tensor, op, group, **kwargs)
|
||||
return tensor
|
||||
|
||||
|
||||
def reduce(tensor, dst, op=dist.ReduceOp.SUM, group=None, **kwargs):
|
||||
if we.is_distributed:
|
||||
return dist.reduce(tensor, dst, op, group, **kwargs)
|
||||
dist.reduce(tensor, dst, op, group, **kwargs)
|
||||
return tensor
|
||||
|
||||
|
||||
def _serialize_to_tensor(data):
|
||||
@@ -243,6 +339,24 @@ def _unserialize_from_tensor(recv_data):
|
||||
return pickle.loads(buffer)
|
||||
|
||||
|
||||
def find_free_port():
|
||||
# Copied from https://github.com/facebookresearch/detectron2/blob/main/detectron2/engine/launch.py # noqa: E501
|
||||
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
||||
# Binding to port 0 will cause the OS to find an available port for us
|
||||
sock.bind(('', 0))
|
||||
port = sock.getsockname()[1]
|
||||
sock.close()
|
||||
# NOTE: there is still a chance the port could be taken by other processes.
|
||||
return port
|
||||
|
||||
|
||||
def is_free_port(port):
|
||||
ips = socket.gethostbyname_ex(socket.gethostname())[-1]
|
||||
ips.append('localhost')
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||
return all(s.connect_ex((ip, port)) != 0 for ip in ips)
|
||||
|
||||
|
||||
def send(tensor, dst, group=None, **kwargs):
|
||||
if we.is_distributed:
|
||||
assert tensor.is_contiguous(
|
||||
@@ -288,6 +402,144 @@ def shared_random_seed():
|
||||
return all_seeds[0]
|
||||
|
||||
|
||||
def all_to_all(x, scatter_dim, gather_dim, group=None, **kwargs):
|
||||
"""
|
||||
`scatter` along one dimension and `gather` along another.
|
||||
"""
|
||||
world_size = dist.get_world_size(group) if we.is_distributed else 1
|
||||
if world_size > 1:
|
||||
inputs = [u.contiguous() for u in x.chunk(world_size, dim=scatter_dim)]
|
||||
outputs = [torch.empty_like(u) for u in inputs]
|
||||
dist.all_to_all(outputs, inputs, group=group, **kwargs)
|
||||
x = torch.cat(outputs, dim=gather_dim).contiguous()
|
||||
return x
|
||||
|
||||
|
||||
def _split(input, dim, group):
|
||||
# skip if world_size == 1
|
||||
rank = dist.get_rank(group=group)
|
||||
world_size = dist.get_world_size(group=group)
|
||||
if world_size == 1:
|
||||
return input
|
||||
|
||||
# split sequence
|
||||
assert input.size(dim) % world_size == 0
|
||||
return input.chunk(world_size, dim=dim)[rank].contiguous()
|
||||
|
||||
|
||||
def _gather(input, dim, group):
|
||||
# skip if world_size == 1
|
||||
world_size = dist.get_world_size(group=group)
|
||||
if world_size == 1:
|
||||
return input
|
||||
|
||||
# gather sequence
|
||||
output = all_gather(input, uniform_size=True, group=group)
|
||||
return torch.cat(output, dim=dim).contiguous()
|
||||
|
||||
|
||||
class AllToAll(Function):
|
||||
@staticmethod
|
||||
def forward(ctx, input, scatter_dim, gather_dim, group):
|
||||
ctx.scatter_dim = scatter_dim
|
||||
ctx.gather_dim = gather_dim
|
||||
ctx.group = group
|
||||
return all_to_all(input, scatter_dim, gather_dim, group)
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
return (all_to_all(grad_output, ctx.gather_dim, ctx.scatter_dim,
|
||||
ctx.group), None, None, None)
|
||||
|
||||
|
||||
class GradScaler(Function):
|
||||
@staticmethod
|
||||
def forward(ctx, input, scale):
|
||||
ctx.scale = scale
|
||||
return input
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
if ctx.scale != 1:
|
||||
grad_output = grad_output * ctx.scale
|
||||
return grad_output, None
|
||||
|
||||
|
||||
class AllGather(Function):
|
||||
@staticmethod
|
||||
def forward(ctx, input, dim, group=None):
|
||||
ctx.dim = dim
|
||||
ctx.group = group
|
||||
output = all_gather(input, uniform_size=True, group=group)
|
||||
return torch.cat(output, dim=dim).contiguous()
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
rank = dist.get_rank(group=ctx.group)
|
||||
world_size = dist.get_world_size(group=ctx.group)
|
||||
return grad_output.chunk(world_size,
|
||||
dim=ctx.dim)[rank].contiguous(), None, None
|
||||
|
||||
|
||||
def diff_all_to_all(input, scatter_dim, gather_dim, group=None):
|
||||
return AllToAll.apply(input, scatter_dim, gather_dim, group)
|
||||
|
||||
|
||||
def diff_scatter_sequence(input, dim, group=None):
|
||||
rank = dist.get_rank(group)
|
||||
world_size = dist.get_world_size(group)
|
||||
output = input.chunk(world_size, dim=dim)[rank].contiguous()
|
||||
return GradScaler.apply(output, 1. / world_size)
|
||||
|
||||
|
||||
def diff_gather_sequence(input, dim, group=None):
|
||||
world_size = dist.get_world_size(group)
|
||||
output = AllGather.apply(input, dim, group)
|
||||
return GradScaler.apply(output, world_size)
|
||||
|
||||
|
||||
class SplitForwardGatherBackward(Function):
|
||||
@staticmethod
|
||||
def forward(ctx, input, dim, group=None, grad_scale=None):
|
||||
ctx.dim = dim
|
||||
ctx.group = group
|
||||
ctx.grad_scale = grad_scale
|
||||
return _split(input, dim, group)
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
if ctx.grad_scale == 'up':
|
||||
grad_output = grad_output * dist.get_world_size(group=ctx.group)
|
||||
elif ctx.grad_scale == 'down':
|
||||
grad_output = grad_output / dist.get_world_size(group=ctx.group)
|
||||
return _gather(grad_output, ctx.dim, ctx.group), None, None, None
|
||||
|
||||
|
||||
class GatherForwardSplitBackward(Function):
|
||||
@staticmethod
|
||||
def forward(ctx, input, dim, group=None, grad_scale=None):
|
||||
ctx.dim = dim
|
||||
ctx.group = group
|
||||
ctx.grad_scale = grad_scale
|
||||
return _gather(input, dim, group)
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
if ctx.grad_scale == 'up':
|
||||
grad_output = grad_output * dist.get_world_size(group=ctx.group)
|
||||
elif ctx.grad_scale == 'down':
|
||||
grad_output = grad_output / dist.get_world_size(group=ctx.group)
|
||||
return _split(grad_output, ctx.dim, ctx.group), None, None, None
|
||||
|
||||
|
||||
def split_forward_gather_backward(input, dim, group=None, grad_scale=None):
|
||||
return SplitForwardGatherBackward.apply(input, dim, group, grad_scale)
|
||||
|
||||
|
||||
def gather_forward_split_backward(input, dim, group=None, grad_scale=None):
|
||||
return GatherForwardSplitBackward.apply(input, dim, group, grad_scale)
|
||||
|
||||
|
||||
global we
|
||||
|
||||
|
||||
@@ -295,7 +547,10 @@ def mp_worker(gpu, ngpus_per_node, cfg, fn, pmi_rank, world_size, work_env):
|
||||
rank = pmi_rank * ngpus_per_node + gpu
|
||||
work_env.device_id = gpu % ngpus_per_node
|
||||
work_env.rank = rank
|
||||
dist.init_process_group(backend='nccl', world_size=world_size, rank=rank)
|
||||
dist.init_process_group(backend='nccl',
|
||||
world_size=world_size,
|
||||
rank=rank,
|
||||
timeout=timedelta(seconds=18000))
|
||||
torch.backends.cudnn.deterministic = cfg.ENV.get('CUDNN_DETERMINISTIC',
|
||||
True)
|
||||
torch.backends.cudnn.benchmark = cfg.ENV.get('CUDNN_BENCHMARK', False)
|
||||
@@ -307,11 +562,50 @@ def mp_worker(gpu, ngpus_per_node, cfg, fn, pmi_rank, world_size, work_env):
|
||||
work_env.logger.info(f'PMI rank {pmi_rank}!')
|
||||
work_env.logger.info(f'Nums of gpu {ngpus_per_node}!')
|
||||
work_env.logger.info(
|
||||
f'Current rank {work_env.rank} current devices num {ngpus_per_node} '
|
||||
f'Current rank {work_env.rank} current devices num {ngpus_per_node} \n'
|
||||
f'current machine rank {pmi_rank} and all world size {world_size}')
|
||||
# model parallel
|
||||
tensor_parallel_size = cfg.ENV.get('TENSOR_PARALLEL_SIZE', 1)
|
||||
pipeline_parallel_size = cfg.ENV.get('PIPELINE_PARALLEL_SIZE', 1)
|
||||
|
||||
we.set_env(work_env.get_env())
|
||||
env_dict = work_env.get_env()
|
||||
|
||||
if tensor_parallel_size * pipeline_parallel_size > 1:
|
||||
'''
|
||||
'''
|
||||
assert world_size % tensor_parallel_size == 0
|
||||
assert world_size % (tensor_parallel_size *
|
||||
pipeline_parallel_size) == 0
|
||||
data_parallel_size = world_size // (tensor_parallel_size *
|
||||
pipeline_parallel_size)
|
||||
mesh = torch.arange(world_size).view(data_parallel_size,
|
||||
pipeline_parallel_size,
|
||||
tensor_parallel_size)
|
||||
index = torch.where(mesh == rank)
|
||||
assert all(u.numel() == 1 for u in index)
|
||||
index = [u.item() for u in index]
|
||||
for j in range(pipeline_parallel_size):
|
||||
for k in range(tensor_parallel_size):
|
||||
group = dist.new_group(mesh[:, j, k].tolist())
|
||||
if j == index[1] and k == index[2]:
|
||||
env_dict['data_parallel_group'] = group
|
||||
for i in range(data_parallel_size):
|
||||
for j in range(pipeline_parallel_size):
|
||||
group = dist.new_group(mesh[i, j, :].tolist())
|
||||
if i == index[0] and j == index[1]:
|
||||
env_dict['tensor_parallel_group'] = group
|
||||
for i in range(data_parallel_size):
|
||||
for k in range(tensor_parallel_size):
|
||||
ranks = mesh[i, :, k].tolist()
|
||||
group = dist.new_group(ranks)
|
||||
if i == index[0] and k == index[2]:
|
||||
env_dict['pipeline_parallel_group'] = group
|
||||
env_dict['pipeline_parallel_ranks'] = ranks
|
||||
we.set_env(env_dict)
|
||||
work_env.logger.info(str(we))
|
||||
fn(cfg)
|
||||
torch.cuda.synchronize()
|
||||
barrier()
|
||||
|
||||
|
||||
class Workenv(object):
|
||||
@@ -322,6 +616,7 @@ class Workenv(object):
|
||||
self.rank = 0
|
||||
self.world_size = 1
|
||||
self.device_id = 0
|
||||
self.backend = ''
|
||||
self.device_count = 1
|
||||
self.seed = 2023
|
||||
self.debug = False
|
||||
@@ -330,24 +625,36 @@ class Workenv(object):
|
||||
self.data_online = False
|
||||
self.share_storage = False
|
||||
|
||||
self.data_parallel_group = None
|
||||
self.tensor_parallel_group = None
|
||||
self.pipleline_parallel_group = None
|
||||
self.pipeline_parallel_ranks = None
|
||||
|
||||
def init_env(self, config, fn, logger=None):
|
||||
# if use pytorch_lightning: then direct use pytorch_lightning.
|
||||
config.ENV = config.get('ENV', {})
|
||||
self.seed = config.ENV.get('SEED', 2023)
|
||||
self.debug = os.environ.get('ES_DEBUG', None) == 'true'
|
||||
|
||||
if logger is None:
|
||||
self.logger = StdMsg(name='env')
|
||||
else:
|
||||
self.logger = logger
|
||||
|
||||
self.sys_envs = config.ENV.get('SYS_ENVS', None)
|
||||
if self.sys_envs:
|
||||
for k, v in self.sys_envs.items():
|
||||
os.environ[k] = v
|
||||
self.logger.info(f'Set env variable {k}={v}')
|
||||
set_random_seed(self.seed)
|
||||
if logger is not None:
|
||||
logger.info(f'And running with seed {self.seed}!')
|
||||
self.logger.info(f'And running with seed {self.seed}!')
|
||||
|
||||
if config.ENV.get('USE_PL', False):
|
||||
self.use_pl = config.ENV.USE_PL
|
||||
fn(config)
|
||||
return
|
||||
if hasattr(config, 'args') and hasattr(config.args, 'launcher'):
|
||||
self.launcher = config.args.launcher
|
||||
if logger is None:
|
||||
self.logger = StdMsg(name='env')
|
||||
else:
|
||||
self.logger = logger
|
||||
|
||||
self.data_online = os.environ.get('DATA_ONLINE', None) == 'true'
|
||||
self.share_storage = os.environ.get('SHARE_STORAGE', None) == 'true'
|
||||
@@ -379,7 +686,8 @@ class Workenv(object):
|
||||
if self.is_distributed:
|
||||
self.backend = config.ENV.get('BACKEND', 'nccl')
|
||||
self.sync_bn = config.ENV.get('SYNC_BN', False)
|
||||
dist.init_process_group(backend=self.backend)
|
||||
dist.init_process_group(backend=self.backend,
|
||||
timeout=timedelta(seconds=18000))
|
||||
# dist.barrier()
|
||||
self.initialized = True
|
||||
if dist.is_initialized():
|
||||
@@ -422,7 +730,7 @@ class Workenv(object):
|
||||
self.device_count = ngpus_per_node
|
||||
world_size = ngpus_per_node * pmi_world_size
|
||||
self.world_size = world_size
|
||||
if self.world_size > 1:
|
||||
if self.world_size >= 1:
|
||||
self.is_distributed = True
|
||||
self.initialized = True
|
||||
if self.is_distributed:
|
||||
@@ -434,13 +742,41 @@ class Workenv(object):
|
||||
self))
|
||||
|
||||
def get_env(self):
|
||||
return self.__dict__
|
||||
ret_dict = {}
|
||||
for k, v in self.__dict__.items():
|
||||
if isinstance(v, (list, dict, int, float, str, bool)):
|
||||
ret_dict[k] = v
|
||||
return ret_dict
|
||||
|
||||
def set_env(self, we_env):
|
||||
for k, v in we_env.items():
|
||||
setattr(self, k, v)
|
||||
set_random_seed(self.seed)
|
||||
|
||||
def group_info(self, group):
|
||||
group_info = f'group size: {group.size()}\n'
|
||||
group_info += f'group rank: {group.rank()}\n'
|
||||
group_info += f'group name: {group.name()}\n'
|
||||
return group_info
|
||||
|
||||
@property
|
||||
def data_group_world_size(self):
|
||||
if self.data_parallel_group is not None:
|
||||
return self.data_parallel_group.size()
|
||||
return self.world_size
|
||||
|
||||
@property
|
||||
def tensor_group_world_size(self):
|
||||
if self.tensor_parallel_group is not None:
|
||||
return self.tensor_parallel_group.size()
|
||||
return 1
|
||||
|
||||
@property
|
||||
def pipeline_group_world_size(self):
|
||||
if self.pipleline_parallel_group is not None:
|
||||
return self.pipleline_parallel_group.size()
|
||||
return 1
|
||||
|
||||
def __str__(self):
|
||||
environ_str = f'Now running in the distributed environment with world size {self.world_size}\n!'
|
||||
environ_str += f'Current pod have {self.device_count} devices!\n'
|
||||
@@ -448,8 +784,19 @@ class Workenv(object):
|
||||
environ_str += f"Current task's global rank is {self.rank} \n"
|
||||
environ_str += f"Current task's data online is set {self.data_online} \n"
|
||||
environ_str += f"Current task's share storage is set {self.share_storage} \n"
|
||||
environ_str += f"Current task's global seed is set {self.seed}"
|
||||
environ_str += f"Current task's global seed is set {self.seed} \n"
|
||||
if self.data_parallel_group is not None:
|
||||
environ_str += f"Current task's data parallel group: {self.group_info(self.data_parallel_group)} \n"
|
||||
if self.pipleline_parallel_group is not None:
|
||||
environ_str += f"Current task's pipeline parallel group: {self.group_info(self.pipleline_parallel_group)} \n"
|
||||
if self.tensor_parallel_group is not None:
|
||||
environ_str += f"Current task's tensor parallel group: {self.group_info(self.tensor_parallel_group)} \n"
|
||||
environ_str += f"Current task's backend is set {self.backend} \n"
|
||||
return environ_str
|
||||
|
||||
def __del__(self):
|
||||
if we.is_distributed:
|
||||
dist.destroy_process_group()
|
||||
|
||||
|
||||
we = Workenv()
|
||||
|
||||
@@ -298,8 +298,7 @@ class AliyunOssFs(BaseFs):
|
||||
try:
|
||||
_ = self._download_object_multi_part(target_path,
|
||||
temp_file,
|
||||
chunk_size=50 *
|
||||
1024 * 1024)
|
||||
chunk_size=size // 100)
|
||||
break
|
||||
except Exception as e:
|
||||
retry += 1
|
||||
@@ -361,7 +360,7 @@ class AliyunOssFs(BaseFs):
|
||||
|
||||
def download_one_part(key):
|
||||
while not slice_queue.empty():
|
||||
R.acquire()
|
||||
R.acquire(timeout=60)
|
||||
try:
|
||||
if not slice_queue.empty():
|
||||
part_number, chunk = slice_queue.get_nowait()
|
||||
@@ -383,7 +382,7 @@ class AliyunOssFs(BaseFs):
|
||||
target_path, chunk[0], chunk[1])
|
||||
with open(temp_part_file, 'wb') as f:
|
||||
f.write(data)
|
||||
R.acquire()
|
||||
R.acquire(timeout=60)
|
||||
try:
|
||||
ret_slice_queue.put_nowait({
|
||||
'part_number':
|
||||
@@ -403,7 +402,7 @@ class AliyunOssFs(BaseFs):
|
||||
'Download part {} for {} error {} retry {} times!'.
|
||||
format(part_number, key, e, retry))
|
||||
if retry >= self._retry_times:
|
||||
R.acquire()
|
||||
R.acquire(timeout=60)
|
||||
try:
|
||||
ret_slice_queue.put_nowait({
|
||||
'part_number': part_number,
|
||||
@@ -539,7 +538,7 @@ class AliyunOssFs(BaseFs):
|
||||
meta_dict[target_path] = etag
|
||||
else:
|
||||
local_path = None
|
||||
R.acquire()
|
||||
R.acquire(timeout=60)
|
||||
try:
|
||||
data_quene.put_nowait([target_path, local_path])
|
||||
except Exception:
|
||||
@@ -623,7 +622,11 @@ class AliyunOssFs(BaseFs):
|
||||
meta_dict = self._get_dir(target_path,
|
||||
local_path=local_path,
|
||||
meta_dict=copy.deepcopy(meta_dict))
|
||||
json.dump(meta_dict, open(check_file, 'w'))
|
||||
for _ in range(5):
|
||||
try:
|
||||
json.dump(meta_dict, open(check_file, 'w'))
|
||||
except:
|
||||
time.sleep(1)
|
||||
if is_tmp:
|
||||
self.add_temp_file(local_path)
|
||||
return local_path
|
||||
@@ -698,7 +701,7 @@ class AliyunOssFs(BaseFs):
|
||||
|
||||
def upload_one_part(key, upload_id):
|
||||
while not slice_queue.empty():
|
||||
R.acquire()
|
||||
R.acquire(timeout=60)
|
||||
try:
|
||||
if not slice_queue.empty():
|
||||
part_number, offset, num_to_upload = slice_queue.get_nowait(
|
||||
@@ -721,7 +724,7 @@ class AliyunOssFs(BaseFs):
|
||||
try:
|
||||
result = _bucket.upload_part(key, upload_id,
|
||||
part_number, raw_data)
|
||||
R.acquire()
|
||||
R.acquire(timeout=60)
|
||||
try:
|
||||
ret_slice_queue.put_nowait({
|
||||
'part_number': part_number,
|
||||
@@ -738,7 +741,7 @@ class AliyunOssFs(BaseFs):
|
||||
'Upload part {} for {} error {} retry {} times!'.
|
||||
format(part_number, key, e, retry))
|
||||
if retry >= self._retry_times:
|
||||
R.acquire()
|
||||
R.acquire(timeout=60)
|
||||
try:
|
||||
ret_slice_queue.put_nowait({
|
||||
'part_number': part_number,
|
||||
@@ -965,8 +968,8 @@ class AliyunOssFs(BaseFs):
|
||||
key,
|
||||
lifecycle,
|
||||
slash_safe=slash_safe)
|
||||
_bucket.put_object_acl(key, oss2.OBJECT_ACL_PUBLIC_READ)
|
||||
if set_public:
|
||||
_bucket.put_object_acl(key, oss2.OBJECT_ACL_PUBLIC_READ)
|
||||
output_url = output_url.replace('%2F', '/').split('?')[0]
|
||||
return output_url
|
||||
except Exception as e:
|
||||
@@ -1060,7 +1063,7 @@ class AliyunOssFs(BaseFs):
|
||||
local_path, target_path)
|
||||
else:
|
||||
flg = False
|
||||
R.acquire()
|
||||
R.acquire(timeout=60)
|
||||
try:
|
||||
data_quene.put_nowait([local_path, target_path, flg])
|
||||
except Exception:
|
||||
|
||||
@@ -39,6 +39,7 @@ class BaseFs(object, metaclass=ABCMeta):
|
||||
self._temp_files = set()
|
||||
self.cfg = cfg
|
||||
self.tmp_dir = cfg.get('TEMP_DIR', None)
|
||||
self.enable_md5_path = cfg.get('ENABLE_MD5_PATH', True)
|
||||
self.auto_clean = cfg.get('AUTO_CLEAN', False)
|
||||
if self.tmp_dir is None:
|
||||
self.auto_clean = True
|
||||
@@ -306,7 +307,10 @@ class BaseFs(object, metaclass=ABCMeta):
|
||||
rand_name += f'{suffix}'
|
||||
tmp_file = osp.join(tempfile.gettempdir(), rand_name)
|
||||
else:
|
||||
cache_name = '{}{}{}'.format(etag, get_md5(target_path), suffix)
|
||||
if self.enable_md5_path:
|
||||
cache_name = '{}{}{}'.format(etag, get_md5(target_path), suffix)
|
||||
else:
|
||||
cache_name = ''
|
||||
tmp_file = osp.join(self.tmp_dir, cache_name)
|
||||
return tmp_file
|
||||
|
||||
|
||||
@@ -84,6 +84,7 @@ class HuggingfaceFs(BaseFs):
|
||||
target_path,
|
||||
local_path=None,
|
||||
wait_finish=False,
|
||||
multi_thread=False,
|
||||
timeout=3600,
|
||||
sign_key=None,
|
||||
worker_id=-1) -> Optional[str]:
|
||||
|
||||
@@ -41,9 +41,7 @@ class LocalFs(BaseFs):
|
||||
if target_path.startswith(self.get_prefix()):
|
||||
return target_path
|
||||
if target_path.startswith('./') or target_path.startswith('../'):
|
||||
return os.path.join(self.get_prefix(),
|
||||
target_path).replace('/./',
|
||||
'/').replace('/../', '/')
|
||||
return os.path.abspath(os.path.join(self.get_prefix(), target_path))
|
||||
if target_path.startswith('/'):
|
||||
return target_path
|
||||
if target_path.startswith('file://'):
|
||||
@@ -236,8 +234,17 @@ class LocalFs(BaseFs):
|
||||
|
||||
def put_object(self, local_data, target_path) -> bool:
|
||||
target_path = self.reconstruct_path(target_path)
|
||||
with open(target_path, 'w') as f:
|
||||
f.write(local_data)
|
||||
dirname = os.path.dirname(target_path)
|
||||
if not os.path.exists(dirname):
|
||||
os.makedirs(dirname, exist_ok=True)
|
||||
if isinstance(local_data, str):
|
||||
with open(target_path, 'w') as f:
|
||||
f.write(local_data)
|
||||
elif isinstance(local_data, bytes):
|
||||
with open(target_path, 'wb') as f:
|
||||
f.write(local_data)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
return True
|
||||
|
||||
def walk_dir(self, file_dir, recurse=True):
|
||||
@@ -249,6 +256,9 @@ class LocalFs(BaseFs):
|
||||
def put_object_from_local_file(self, local_path, target_path) -> bool:
|
||||
target_path = self.reconstruct_path(target_path)
|
||||
local_path = self.reconstruct_path(local_path)
|
||||
dirname = os.path.dirname(target_path)
|
||||
if not os.path.exists(dirname):
|
||||
os.makedirs(dirname, exist_ok=True)
|
||||
if local_path != target_path:
|
||||
try:
|
||||
shutil.copy(local_path, target_path)
|
||||
@@ -271,7 +281,7 @@ class LocalFs(BaseFs):
|
||||
return False
|
||||
return True
|
||||
try:
|
||||
os.makedirs(target_dir)
|
||||
os.makedirs(target_dir, exist_ok=True)
|
||||
except Exception as e:
|
||||
self.logger.error(e)
|
||||
return False
|
||||
@@ -309,6 +319,9 @@ class LocalFs(BaseFs):
|
||||
multi_thread=False) -> bool:
|
||||
local_dir = self.reconstruct_path(local_dir)
|
||||
target_dir = self.reconstruct_path(target_dir)
|
||||
dirname = os.path.dirname(target_dir)
|
||||
if not os.path.exists(dirname):
|
||||
os.makedirs(dirname, exist_ok=True)
|
||||
if local_dir == target_dir:
|
||||
return True
|
||||
# # cp -f local_dir/* target_dir/*
|
||||
|
||||
@@ -24,8 +24,8 @@ class ModelscopeFs(BaseFs):
|
||||
super(ModelscopeFs, self).__init__(cfg, logger=logger)
|
||||
retry_times = cfg.get('RETRY_TIMES', 10)
|
||||
self._retry_times = retry_times
|
||||
self._model_id_loaded = set()
|
||||
self._model_file_loaded = set()
|
||||
self._model_id_loaded = dict()
|
||||
self._model_file_loaded = dict()
|
||||
|
||||
def get_prefix(self) -> str:
|
||||
return 'ms://'
|
||||
@@ -69,14 +69,14 @@ class ModelscopeFs(BaseFs):
|
||||
if revision is not None:
|
||||
local_path = local_path + '_' + str(revision)
|
||||
|
||||
model_file = os.path.join(key, file_path)
|
||||
retry = 0
|
||||
while retry < self._retry_times:
|
||||
try:
|
||||
model_file = os.path.join(key, file_path)
|
||||
if model_file in self._model_file_loaded:
|
||||
local_path = os.path.join(local_path, model_file)
|
||||
local_path = self._model_file_loaded[model_file]
|
||||
if not osp.exists(local_path):
|
||||
self._model_file_loaded.remove(key)
|
||||
self._model_file_loaded.pop(model_file)
|
||||
else:
|
||||
if token is not None:
|
||||
cookies = self.get_modelscope_cookie(token)
|
||||
@@ -94,7 +94,7 @@ class ModelscopeFs(BaseFs):
|
||||
|
||||
if retry >= self._retry_times:
|
||||
return None
|
||||
self._model_file_loaded.add(model_file)
|
||||
self._model_file_loaded[model_file] = local_path
|
||||
if is_tmp:
|
||||
self.add_temp_file(local_path)
|
||||
return local_path
|
||||
@@ -142,9 +142,9 @@ class ModelscopeFs(BaseFs):
|
||||
while retry < self._retry_times:
|
||||
try:
|
||||
if key in self._model_id_loaded:
|
||||
local_path = os.path.join(local_path, key)
|
||||
local_path = self._model_id_loaded[key]
|
||||
if not osp.exists(local_path):
|
||||
self._model_id_loaded.remove(key)
|
||||
self._model_id_loaded.pop(key)
|
||||
else:
|
||||
if token is not None:
|
||||
cookies = self.get_modelscope_cookie(token)
|
||||
@@ -162,7 +162,7 @@ class ModelscopeFs(BaseFs):
|
||||
if retry >= self._retry_times:
|
||||
return None
|
||||
|
||||
self._model_id_loaded.add(key)
|
||||
self._model_id_loaded[key] = local_path
|
||||
if is_tmp:
|
||||
self.add_temp_file(local_path)
|
||||
if not ret_folder == '':
|
||||
|
||||
@@ -281,7 +281,7 @@ class FileSystem(object):
|
||||
else:
|
||||
return False
|
||||
|
||||
def get_batch_objects_from(self, target_path_list, wait_finish=False):
|
||||
def get_batch_objects_from(self, target_path_list, wait_finish=False, return_target_path=False):
|
||||
data_quene = Queue()
|
||||
batch_size = 20
|
||||
R = threading.Lock()
|
||||
@@ -293,7 +293,7 @@ class FileSystem(object):
|
||||
wait_finish=wait_finish)
|
||||
else:
|
||||
local_path = None
|
||||
R.acquire()
|
||||
R.acquire(timeout = 2)
|
||||
try:
|
||||
data_quene.put_nowait([target_path, local_path])
|
||||
except Exception:
|
||||
@@ -320,7 +320,10 @@ class FileSystem(object):
|
||||
|
||||
for target_path in batch_list:
|
||||
local_path = file_dict.get(target_path, None)
|
||||
yield local_path
|
||||
if return_target_path:
|
||||
yield target_path, local_path
|
||||
else:
|
||||
yield local_path
|
||||
|
||||
def put_batch_objects_to(self,
|
||||
local_path_list,
|
||||
@@ -352,7 +355,7 @@ class FileSystem(object):
|
||||
pass
|
||||
else:
|
||||
flg = False
|
||||
R.acquire()
|
||||
R.acquire(timeout=2)
|
||||
try:
|
||||
data_quene.put_nowait([local_path, target_path, flg])
|
||||
except Exception:
|
||||
|
||||
@@ -6,7 +6,6 @@ from tqdm import tqdm
|
||||
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
|
||||
def init_1level_llfs(list_file,
|
||||
max_lines=1024,
|
||||
index_name='index',
|
||||
@@ -14,6 +13,57 @@ def init_1level_llfs(list_file,
|
||||
r"""Construct large-list-file index.
|
||||
"""
|
||||
index_dir = osp.splitext(list_file)[0]
|
||||
total_failed = 0
|
||||
with FS.get_from(list_file, wait_finish=True) as local_path:
|
||||
_num_split, index_stack, save_files = 0, [], []
|
||||
cache_data_list, target_path_list = [], []
|
||||
with open(local_path, 'r', buffering=1000000) as f:
|
||||
for line in tqdm(f):
|
||||
index_stack.append(line.strip())
|
||||
if len(index_stack) >= max_lines:
|
||||
save_file = f'{index_dir}/{index_name}/{_num_split + 1:09d}.txt'
|
||||
cache_data_list.append('\n'.join(index_stack).encode())
|
||||
target_path_list.append(save_file)
|
||||
# with FS.put_to(save_file) as cache_path:
|
||||
# with open(cache_path, 'w') as f_w:
|
||||
# f_w.write('\n'.join(index_stack))
|
||||
if len(target_path_list) >= 1000:
|
||||
put_res = [(target_path, flg) for local_path, target_path, flg
|
||||
in FS.put_batch_objects_to(cache_data_list, target_path_list, batch_size=40)]
|
||||
for put_flg in put_res:
|
||||
if not put_flg:
|
||||
total_failed += 1
|
||||
cache_data_list, target_path_list = [], []
|
||||
index_stack = []
|
||||
_num_split += 1
|
||||
save_files.append(save_file)
|
||||
put_res = [(target_path, flg) for local_path, target_path, flg
|
||||
in FS.put_batch_objects_to(cache_data_list, target_path_list, batch_size=50)]
|
||||
for put_flg in put_res:
|
||||
if not put_flg:
|
||||
total_failed += 1
|
||||
print(f'Failed to put {total_failed} files.')
|
||||
if len(index_stack) > 0:
|
||||
save_file = f'{index_dir}/{index_name}/{_num_split + 1:06d}.txt'
|
||||
with FS.put_to(save_file) as cache_path:
|
||||
with open(cache_path, 'w') as f_w:
|
||||
f_w.write('\n'.join(index_stack))
|
||||
save_files.append(save_file)
|
||||
|
||||
# output meta-file
|
||||
index_file = osp.join(index_dir, f'{index_name}.txt')
|
||||
with FS.put_to(index_file) as cache_path:
|
||||
with open(cache_path, 'w') as f_w:
|
||||
f_w.write('\n'.join(save_files))
|
||||
return index_file
|
||||
|
||||
def init_1level_llfs_single_threading(list_file,
|
||||
max_lines=1024,
|
||||
index_name='index',
|
||||
delimiter='\n'):
|
||||
r"""Construct large-list-file index.
|
||||
"""
|
||||
index_dir = osp.splitext(list_file)[0]
|
||||
print(list_file)
|
||||
with FS.get_from(list_file, wait_finish=True) as local_path:
|
||||
print(local_path)
|
||||
|
||||
@@ -47,7 +47,7 @@ def time_since(since, percent):
|
||||
return '{} {:.2f}%({})'.format(as_time(s), 100 * percent, as_time(rs))
|
||||
|
||||
|
||||
def get_logger(name='torch dist'):
|
||||
def get_logger(name='scepter', level=logging.INFO):
|
||||
logger = logging.getLogger(name)
|
||||
logger.propagate = False
|
||||
if len(logger.handlers) == 0:
|
||||
@@ -57,8 +57,8 @@ def get_logger(name='torch dist'):
|
||||
'[File: %(filename)s Function: %(funcName)s at line %(lineno)d] %(message)s'
|
||||
)
|
||||
std_handler.setFormatter(formatter)
|
||||
std_handler.setLevel(logging.INFO)
|
||||
logger.setLevel(logging.INFO)
|
||||
std_handler.setLevel(level)
|
||||
logger.setLevel(level)
|
||||
logger.addHandler(std_handler)
|
||||
return logger
|
||||
|
||||
|
||||
+328
-123
@@ -1,7 +1,9 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import json
|
||||
import os.path
|
||||
from io import BytesIO
|
||||
from numbers import Number
|
||||
|
||||
import numpy as np
|
||||
@@ -76,7 +78,8 @@ def merge_gathered_probe(all_gathered_data):
|
||||
is_image=ret_data.is_image,
|
||||
build_html=ret_data.build_html,
|
||||
build_label=ret_data.build_label,
|
||||
view_distribute=ret_data.view_distribute)
|
||||
view_distribute=ret_data.view_distribute,
|
||||
is_presave=ret_data.is_presave)
|
||||
elif isinstance(ret_data.data, list):
|
||||
if ret_data.build_label is not None:
|
||||
if isinstance(ret_data.build_label, str):
|
||||
@@ -95,19 +98,57 @@ def merge_gathered_probe(all_gathered_data):
|
||||
is_image=ret_data.is_image,
|
||||
build_html=ret_data.build_html,
|
||||
build_label=ret_data.build_label,
|
||||
view_distribute=ret_data.view_distribute)
|
||||
view_distribute=ret_data.view_distribute,
|
||||
is_presave=ret_data.is_presave)
|
||||
else:
|
||||
all_gathered_data[key] = gathered_data
|
||||
return all_gathered_data
|
||||
|
||||
|
||||
class MediaHandler():
|
||||
def __init__(self, batch_size = 10):
|
||||
self.file_list = []
|
||||
self.target_path_list = []
|
||||
self.target_status = {}
|
||||
self.batch_size = batch_size
|
||||
|
||||
def append(self, source_file, target_path):
|
||||
self.file_list.append(source_file)
|
||||
self.target_path_list.append(target_path)
|
||||
if len(self.file_list) > 2 * self.batch_size:
|
||||
generator = FS.put_batch_objects_to(self.file_list, self.target_path_list, batch_size=self.batch_size)
|
||||
for local_path, target_path, flg in generator:
|
||||
self.target_status[target_path] = flg
|
||||
self.file_list.clear()
|
||||
self.target_path_list.clear()
|
||||
def sync(self):
|
||||
if len(self.file_list) > 0:
|
||||
if len(self.file_list) > 4 * self.batch_size:
|
||||
generator = FS.put_batch_objects_to(self.file_list, self.target_path_list, batch_size=self.batch_size)
|
||||
for local_path, target_path, flg in generator:
|
||||
self.target_status[target_path] = flg
|
||||
else:
|
||||
for file_, target_path in zip(self.file_list, self.target_path_list):
|
||||
self.target_status[target_path] = FS.put_object(file_.getvalue(), target_path)
|
||||
self.file_list.clear()
|
||||
self.target_path_list.clear()
|
||||
|
||||
def clear(self):
|
||||
self.file_list.clear()
|
||||
self.target_path_list.clear()
|
||||
self.target_status.clear()
|
||||
|
||||
|
||||
class ProbeData():
|
||||
def __init__(self,
|
||||
data,
|
||||
is_image=False,
|
||||
is_video=False,
|
||||
fps=8,
|
||||
build_html=False,
|
||||
build_label=None,
|
||||
view_distribute=False):
|
||||
view_distribute=False,
|
||||
is_presave = False):
|
||||
''' Probe Data Initialize.
|
||||
We only support basic types such as [torch.Tensor, numpy.ndarray, number, str],
|
||||
or [dict, list] of [dict, list,
|
||||
@@ -164,6 +205,28 @@ class ProbeData():
|
||||
elif isinstance(v, np.ndarray):
|
||||
data[k] = v
|
||||
self.basic_type = False
|
||||
elif isinstance(v, list):
|
||||
for v_idx, v_v in enumerate(v):
|
||||
if not check_legal_type(v_v):
|
||||
if isinstance(v_v, torch.Tensor):
|
||||
data[k][v_idx] = v_v.detach().cpu().numpy()
|
||||
self.basic_type = False
|
||||
elif isinstance(v_v, np.ndarray):
|
||||
data[k][v_idx] = v_v
|
||||
self.basic_type = False
|
||||
else:
|
||||
raise f'Unsupport data type for {v_v}'
|
||||
elif isinstance(v, dict):
|
||||
for k_k, v_v in v.items():
|
||||
if not check_legal_type(v_v):
|
||||
if isinstance(v_v, torch.Tensor):
|
||||
data[k][k_k] = v_v.detach().cpu().numpy()
|
||||
self.basic_type = False
|
||||
elif isinstance(v_v, np.ndarray):
|
||||
data[k][k_k] = v_v
|
||||
self.basic_type = False
|
||||
else:
|
||||
raise f'Unsupport data type for {v_v}'
|
||||
else:
|
||||
raise f'Unsupport data type for {v}'
|
||||
self.data = data
|
||||
@@ -176,6 +239,28 @@ class ProbeData():
|
||||
elif isinstance(v, np.ndarray):
|
||||
data[idx] = v
|
||||
self.basic_type = False
|
||||
elif isinstance(v, list):
|
||||
for v_idx, v_v in enumerate(v):
|
||||
if not check_legal_type(v_v):
|
||||
if isinstance(v_v, torch.Tensor):
|
||||
data[idx][v_idx] = v_v.detach().cpu().numpy()
|
||||
self.basic_type = False
|
||||
elif isinstance(v_v, np.ndarray):
|
||||
data[idx][v_idx] = v_v
|
||||
self.basic_type = False
|
||||
else:
|
||||
raise f'Unsupport data type for {v_v}'
|
||||
elif isinstance(v, dict):
|
||||
for k, v_v in v.items():
|
||||
if not check_legal_type(v_v):
|
||||
if isinstance(v_v, torch.Tensor):
|
||||
data[idx][k] = v_v.detach().cpu().numpy()
|
||||
self.basic_type = False
|
||||
elif isinstance(v_v, np.ndarray):
|
||||
data[idx][k] = v_v
|
||||
self.basic_type = False
|
||||
else:
|
||||
raise f'Unsupport data type for {v_v}'
|
||||
else:
|
||||
raise f'Unsupport data type for {v}'
|
||||
self.data = data
|
||||
@@ -185,7 +270,13 @@ class ProbeData():
|
||||
raise f'Unsupport data type for {data}'
|
||||
|
||||
self.is_image = is_image
|
||||
self.is_video = is_video
|
||||
self.image_postfix = 'jpg'
|
||||
self.video_postfix = 'mp4'
|
||||
self.fps = fps
|
||||
self.build_html = build_html
|
||||
self.media_handler = MediaHandler()
|
||||
self.is_presave = is_presave
|
||||
|
||||
if self.build_html:
|
||||
assert build_label is not None
|
||||
@@ -199,7 +290,84 @@ class ProbeData():
|
||||
build_label, dict)
|
||||
self.build_label = build_label
|
||||
|
||||
def save_image(self, file_prefix, images, image_postfix):
|
||||
def get_format(self, extension):
|
||||
if extension.lower() in ['jpg', 'jpeg']:
|
||||
return 'JPEG'
|
||||
if extension.lower() in ['png']:
|
||||
return 'PNG'
|
||||
return 'JPEG'
|
||||
def save_one_video(self, file_path, videos, fps = 8):
|
||||
# write video
|
||||
import imageio
|
||||
try:
|
||||
writer = imageio.get_writer(file_path, fps=fps, format=".mp4", codec='libx264', quality=8)
|
||||
for frame in videos:
|
||||
writer.append_data(frame)
|
||||
writer.close()
|
||||
return True
|
||||
except:
|
||||
return False
|
||||
|
||||
|
||||
def save_video(self, file_prefix, videos, video_postfix, fps = 8, rank = 0):
|
||||
if isinstance(videos, list):
|
||||
for video in videos:
|
||||
if isinstance(video, list):
|
||||
raise f"Only surpport one layer nested list."
|
||||
return [self.save_video(file_prefix + f'_{rank}_{idx}', v, video_postfix, fps) for idx, v in enumerate(videos)]
|
||||
np_shape = videos.shape
|
||||
# 4D
|
||||
shape_str = '_'.join([str(v) for v in np_shape])
|
||||
if len(np_shape) == 5:
|
||||
# channel is 1 or 3
|
||||
if np_shape[-1] == 1 or np_shape[-1] == 3 or np_shape[-1] == 4:
|
||||
if np_shape[-1] == 1:
|
||||
videos = videos.reshape(videos.shape[:-1])
|
||||
file_list = []
|
||||
for idx in range(np_shape[0]):
|
||||
if videos[idx].shape[0] > 1:
|
||||
file_path = os.path.join(file_prefix, f'probe_{rank}_{idx}_[{shape_str}].{video_postfix}')
|
||||
byio = BytesIO()
|
||||
is_suc = self.save_one_video(byio, videos[idx], fps)
|
||||
if not is_suc:
|
||||
byio.write(b"")
|
||||
else:
|
||||
file_path = os.path.join(file_prefix, f'probe_{rank}_{idx}_[{shape_str}].{self.image_postfix}')
|
||||
byio = BytesIO()
|
||||
Image.fromarray(videos[idx][0]).save(byio, self.get_format(self.image_postfix))
|
||||
self.media_handler.append(byio, file_path)
|
||||
file_list.append(file_path)
|
||||
return file_list
|
||||
else:
|
||||
raise f"Ensure your data's dim is BWHC, and channel is 1 or 3 for {file_prefix}"
|
||||
elif len(np_shape) == 4:
|
||||
if np_shape[-1] == 1 or np_shape[-1] == 3 or np_shape[-1] == 4:
|
||||
if np_shape[-1] == 1:
|
||||
videos = videos.reshape(videos.shape[:-1])
|
||||
if videos.shape[0] > 1:
|
||||
file_path = file_prefix + f'_probe_{rank}_[{shape_str}].{video_postfix}'
|
||||
byio = BytesIO()
|
||||
is_suc = self.save_one_video(byio, videos, fps)
|
||||
if not is_suc:
|
||||
byio.write(b"")
|
||||
else:
|
||||
file_path = file_prefix + f'_probe_{rank}_[{shape_str}].{self.image_postfix}'
|
||||
byio = BytesIO()
|
||||
Image.fromarray(videos[0]).save(byio, self.get_format(self.image_postfix))
|
||||
self.media_handler.append(byio, file_path)
|
||||
return file_path
|
||||
else:
|
||||
videos = videos.reshape(list(videos.shape) + [1])
|
||||
return self.save_video(file_prefix, videos, video_postfix, fps = fps)
|
||||
else:
|
||||
raise f"Ensure your data's dim is BFWHC or FWHC, and channel is 1 or 3 for {file_prefix}"
|
||||
|
||||
def save_image(self, file_prefix, images, image_postfix, rank = 0):
|
||||
if isinstance(images, list):
|
||||
for image in images:
|
||||
if isinstance(image, list):
|
||||
raise f"Only surpport one layer nested list."
|
||||
return [self.save_image(file_prefix + f'_{rank}_{idx}', v, image_postfix) for idx, v in enumerate(images)]
|
||||
np_shape = images.shape
|
||||
# 4D
|
||||
shape_str = '_'.join([str(v) for v in np_shape])
|
||||
@@ -210,9 +378,10 @@ class ProbeData():
|
||||
images = images.reshape(images.shape[:-1])
|
||||
file_list = []
|
||||
for idx in range(np_shape[0]):
|
||||
file_path = file_prefix + f'_probe_{idx}_[{shape_str}].{image_postfix}'
|
||||
with FS.put_to(file_path) as local_path:
|
||||
Image.fromarray(images[idx, ...]).save(local_path)
|
||||
file_path = os.path.join(file_prefix, f'probe_{rank}_{idx}_[{shape_str}].{image_postfix}')
|
||||
byio = BytesIO()
|
||||
Image.fromarray(images[idx, ...]).save(byio, self.get_format(self.image_postfix))
|
||||
self.media_handler.append(byio, file_path)
|
||||
file_list.append(file_path)
|
||||
return file_list
|
||||
else:
|
||||
@@ -221,26 +390,29 @@ class ProbeData():
|
||||
if np_shape[-1] == 1 or np_shape[-1] == 3 or np_shape[-1] == 4:
|
||||
if np_shape[-1] == 1:
|
||||
images = images.reshape(images.shape[:-1])
|
||||
file_path = file_prefix + f'_probe_[{shape_str}].{image_postfix}'
|
||||
with FS.put_to(file_path) as local_path:
|
||||
Image.fromarray(images).save(local_path)
|
||||
file_path = file_prefix + f'_probe_{rank}_[{shape_str}].{image_postfix}'
|
||||
byio = BytesIO()
|
||||
Image.fromarray(images).save(byio, self.get_format(self.image_postfix))
|
||||
self.media_handler.append(byio, file_path)
|
||||
return file_path
|
||||
else:
|
||||
images = images.reshape(list(images.shape) + [1])
|
||||
return self.save_image(file_prefix, images, image_postfix)
|
||||
elif len(np_shape) == 2:
|
||||
file_path = file_prefix + f'_probe_[{shape_str}].{image_postfix}'
|
||||
with FS.put_to(file_path) as local_path:
|
||||
Image.fromarray(images).save(local_path)
|
||||
file_path = file_prefix + f'_probe_{rank}_[{shape_str}].{image_postfix}'
|
||||
byio = BytesIO()
|
||||
Image.fromarray(images).save(byio, self.get_format(self.image_postfix))
|
||||
self.media_handler.append(byio, file_path)
|
||||
return file_path
|
||||
else:
|
||||
raise f"Ensure your data's dim is BWHC or WHC or WH, and channel is 1 or 3 for {file_prefix}"
|
||||
|
||||
def save_npy(self, file_prefix, data):
|
||||
def save_npy(self, file_prefix, data, rank = 0):
|
||||
shape_str = '_'.join([str(v) for v in data.shape])
|
||||
file_path = file_prefix + f'_{shape_str}.npy'
|
||||
with FS.put_to(file_path) as local_path:
|
||||
np.save(local_path, data)
|
||||
file_path = file_prefix + f'_{rank}_{shape_str}.npy'
|
||||
byio = BytesIO()
|
||||
np.save(byio, data)
|
||||
self.media_handler.append(byio, file_path)
|
||||
return file_path
|
||||
|
||||
def save_html(self, html_prefix, ret_data, ret_label):
|
||||
@@ -249,9 +421,10 @@ class ProbeData():
|
||||
with open(local_path, 'w') as f:
|
||||
f.writelines('<meta charset="utf-8">\n')
|
||||
f.writelines('<style>input{height:' + f'{height}px;' +
|
||||
'opacity:1.0;}</style>\n')
|
||||
'opacity:1.0;} textarea {font-size: 32px;}</style>\n')
|
||||
f.writelines('<br><hr/>\n')
|
||||
all_ranks = list()
|
||||
is_textarea = False
|
||||
for save_id, save_data in enumerate(zip(ret_data, ret_label)):
|
||||
save_path, save_label = save_data
|
||||
one_rank = '<table><tr>'
|
||||
@@ -259,15 +432,32 @@ class ProbeData():
|
||||
one_path, one_label = one_data
|
||||
one_label = one_label.replace('<', '<').replace(
|
||||
'>', '>')
|
||||
url = FS.get_url(one_path,
|
||||
lifecycle=3600 * 365 * 24).replace(
|
||||
'.oss-internal.aliyun-inc.',
|
||||
'.oss.aliyuncs.').replace(
|
||||
'-internal', '')
|
||||
one_rank += (
|
||||
f'<td align="center"><input type="image" src="{url}" >'
|
||||
f'<br><font size="4"><strong>{save_id}-{idx}|{one_label}<strong></font><br/></td>'
|
||||
)
|
||||
try:
|
||||
url = FS.get_url(one_path,
|
||||
lifecycle=3600 * 365 * 24).replace(
|
||||
'.oss-internal.aliyun-inc.',
|
||||
'.oss.aliyuncs.').replace(
|
||||
'-internal', '')
|
||||
except:
|
||||
url = one_path
|
||||
if len(one_label) > 10 and idx == len(save_path) - 1:
|
||||
is_textarea = True
|
||||
if self.is_video and one_path.endswith(self.video_postfix):
|
||||
one_rank += f'<td align="center"><video height="{height}" controls="">'
|
||||
one_rank += f'<source src="{url}" type="video/mp4"></video>'
|
||||
if idx == len(save_path) - 1 and is_textarea:
|
||||
one_rank += f'<br><font size="4"><strong>{save_id}-{idx}<strong></font><br/></td>'
|
||||
one_rank += f'<td align="center"><textarea rows="16" cols="40">{one_label}</textarea><br><font size="4"><strong>{save_id}-{idx}<strong></font><br/></td>'
|
||||
else:
|
||||
one_rank += f'<br><font size="4"><strong>{save_id}-{idx}|{one_label}<strong></font><br/></td>'
|
||||
else:
|
||||
one_rank += f'<td align="center"><input type="image" src="{url}" >'
|
||||
if idx == len(save_path) - 1 and is_textarea:
|
||||
one_rank += f'<br><font size="4"><strong>{save_id}-{idx}<strong></font><br/></td>'
|
||||
one_rank += f'<td align="center"><textarea rows="16" cols="40">{one_label}</textarea><br><font size="4"><strong>{save_id}-{idx}<strong></font><br/></td>'
|
||||
else:
|
||||
one_rank += f'<br><font size="4"><strong>{save_id}-{idx}|{one_label}<strong></font><br/></td>'
|
||||
|
||||
one_rank += '</tr></table><hr/>'
|
||||
all_ranks.append(one_rank)
|
||||
f.writelines('\n'.join(all_ranks))
|
||||
@@ -277,117 +467,132 @@ class ProbeData():
|
||||
def distribute(self):
|
||||
return self._distribute_dict
|
||||
|
||||
def to_log(self, prefix=None, image_postfix='jpg'):
|
||||
def save_one_media(self, idx, v, prefix_path, image_postfix, video_postfix, rank = 0):
|
||||
ret_label = None
|
||||
if self.is_image:
|
||||
ret_medias = self.save_image(prefix_path, v,
|
||||
image_postfix, rank = rank)
|
||||
elif self.is_video:
|
||||
ret_medias = self.save_video(prefix_path, v,
|
||||
video_postfix, fps=self.fps, rank = rank)
|
||||
else:
|
||||
ret_data = self.save_npy(prefix_path, v, rank = rank)
|
||||
return ret_data, ret_label
|
||||
ret_data = ret_medias if isinstance(ret_medias, list) else [ret_medias]
|
||||
if self.build_html:
|
||||
if isinstance(ret_medias, list):
|
||||
if isinstance(self.build_label, str):
|
||||
ret_label = [self.build_label for _ in ret_medias]
|
||||
elif isinstance(self.build_label[idx], list):
|
||||
assert len(self.build_label[idx]) == len(
|
||||
ret_medias)
|
||||
ret_label = self.build_label[idx]
|
||||
else:
|
||||
ret_label = [
|
||||
self.build_label[idx]
|
||||
for _ in ret_medias
|
||||
]
|
||||
else:
|
||||
if isinstance(self.build_label, str):
|
||||
ret_label = [self.build_label]
|
||||
else:
|
||||
ret_label = [self.build_label[idx]]
|
||||
return ret_data, ret_label
|
||||
|
||||
def presave(self, prefix=None, image_postfix='jpg', video_postfix='mp4', rank = 0):
|
||||
self.image_postfix = image_postfix
|
||||
self.video_postfix = video_postfix
|
||||
if isinstance(self.data, np.ndarray):
|
||||
if prefix is None:
|
||||
raise 'You should provide the save prefix for array sample.'
|
||||
# save jpg
|
||||
if self.is_image:
|
||||
ret_data = self.save_image(prefix, self.data, image_postfix)
|
||||
if isinstance(ret_data, list):
|
||||
ret_data = [ret_data]
|
||||
if self.build_html:
|
||||
ret_label = []
|
||||
if isinstance(self.build_label, str):
|
||||
ret_label.append(
|
||||
[self.build_label for _ in ret_data[0]])
|
||||
else:
|
||||
ret_label.append(self.build_label)
|
||||
if not len(ret_data[0]) == len(ret_label[0]):
|
||||
raise f"The {prefix} label's length should be equal with 1st dim {self.data.shape[0]}."
|
||||
html_prefix = prefix + '_probe.html'
|
||||
html_file = self.save_html(html_prefix, ret_data,
|
||||
ret_label)
|
||||
return {'ori_file': ret_data, 'html': html_file}
|
||||
else:
|
||||
return {'ori_file': ret_data}
|
||||
else:
|
||||
return ret_data
|
||||
ret_data = self.save_image(prefix, self.data, image_postfix, rank = rank)
|
||||
elif self.is_video:
|
||||
ret_data = self.save_video(prefix, self.data, video_postfix, fps=self.fps, rank = rank)
|
||||
else:
|
||||
ret_data = self.save_npy(prefix, self.data)
|
||||
return ret_data
|
||||
ret_data = self.save_npy(prefix, self.data, rank = rank)
|
||||
self.media_handler.sync()
|
||||
self.media_handler.clear()
|
||||
if isinstance(ret_data, list):
|
||||
ret_data = [ret_data]
|
||||
ret_label = []
|
||||
if self.build_html:
|
||||
if isinstance(self.build_label, str):
|
||||
ret_label.append(
|
||||
[self.build_label for _ in ret_data[0]])
|
||||
else:
|
||||
ret_label.append(self.build_label)
|
||||
if not len(ret_data[0]) == len(ret_label[0]):
|
||||
raise f"The {prefix} label's length should be equal with 1st dim {self.data.shape[0]}."
|
||||
self.data = {"ret_data": ret_data, "ret_label": ret_label}
|
||||
else:
|
||||
self.data = ret_data
|
||||
self.is_presave = True
|
||||
elif isinstance(self.data, list):
|
||||
if not self.basic_type:
|
||||
ret_data = []
|
||||
ret_label = []
|
||||
for idx, v in enumerate(self.data):
|
||||
prefix_path = os.path.join(prefix, f'{idx}')
|
||||
if self.is_image:
|
||||
ret_images = self.save_image(prefix_path, v,
|
||||
image_postfix)
|
||||
ret_data.append(ret_images if isinstance(
|
||||
ret_images, list) else [ret_images])
|
||||
if self.build_html:
|
||||
if isinstance(ret_images, list):
|
||||
if isinstance(self.build_label, str):
|
||||
ret_label.append(
|
||||
[self.build_label for _ in ret_images])
|
||||
elif isinstance(self.build_label[idx], list):
|
||||
assert len(self.build_label[idx]) == len(
|
||||
ret_images)
|
||||
ret_label.append(self.build_label[idx])
|
||||
else:
|
||||
ret_label.append([
|
||||
self.build_label[idx]
|
||||
for _ in ret_images
|
||||
])
|
||||
else:
|
||||
if isinstance(self.build_label, str):
|
||||
ret_label.append([self.build_label])
|
||||
else:
|
||||
ret_label.append([self.build_label[idx]])
|
||||
else:
|
||||
ret_data.append(self.save_npy(prefix_path, v))
|
||||
if self.is_image and self.build_html:
|
||||
html_prefix = prefix + '_probe.html'
|
||||
html_file = self.save_html(html_prefix, ret_data,
|
||||
ret_label)
|
||||
return {'ori_file': ret_data, 'html': html_file}
|
||||
else:
|
||||
return {'ori_file': ret_data}
|
||||
else:
|
||||
return self.data
|
||||
ret_one_data, ret_one_label = self.save_one_media(idx, v,
|
||||
prefix_path,
|
||||
image_postfix,
|
||||
video_postfix,
|
||||
rank=rank)
|
||||
ret_data.append(ret_one_data)
|
||||
ret_label.append(ret_one_label) if ret_one_label is not None else ret_label
|
||||
self.media_handler.sync()
|
||||
self.media_handler.clear()
|
||||
self.data = {"ret_data": ret_data, "ret_label": ret_label}
|
||||
self.is_presave = True
|
||||
elif isinstance(self.data, dict):
|
||||
if not self.basic_type:
|
||||
ret_data = []
|
||||
ret_label = []
|
||||
for k, v in self.data:
|
||||
for k, v in self.data.items():
|
||||
prefix_path = os.path.join(prefix, f'{k}_')
|
||||
if self.is_image:
|
||||
ret_images = self.save_image(prefix_path, v,
|
||||
image_postfix)
|
||||
if isinstance(ret_images, list):
|
||||
ret_data.append(ret_images)
|
||||
else:
|
||||
ret_data.append([ret_images])
|
||||
if self.build_html:
|
||||
if isinstance(ret_images, list):
|
||||
if isinstance(self.build_label, str):
|
||||
ret_label.append(
|
||||
[self.build_label for _ in ret_images])
|
||||
elif isinstance(self.build_label[k], list):
|
||||
assert len(
|
||||
self.build_label[k]) == len(ret_images)
|
||||
ret_label.append(self.build_label[k])
|
||||
else:
|
||||
ret_label.append([
|
||||
self.build_label[k] for _ in ret_images
|
||||
])
|
||||
else:
|
||||
if isinstance(self.build_label, str):
|
||||
ret_label.append([self.build_label])
|
||||
else:
|
||||
ret_label.append([self.build_label[k]])
|
||||
else:
|
||||
ret_data.append(self.save_npy(prefix_path, v))
|
||||
if self.is_image and self.build_html:
|
||||
html_prefix = prefix + '_probe.html'
|
||||
html_file = self.save_html(html_prefix, ret_data,
|
||||
ret_label)
|
||||
return {'ori_file': ret_data, 'html': html_file}
|
||||
else:
|
||||
return {'ori_file': ret_data}
|
||||
else:
|
||||
return self.data
|
||||
else:
|
||||
ret_one_data, ret_one_label = self.save_one_media(k, v,
|
||||
prefix_path,
|
||||
image_postfix,
|
||||
video_postfix,
|
||||
rank = rank)
|
||||
ret_data.append(ret_one_data)
|
||||
ret_label.append(ret_one_label) if ret_one_label is not None else ret_label
|
||||
self.media_handler.sync()
|
||||
self.media_handler.clear()
|
||||
self.data = {"ret_data": ret_data, "ret_label": ret_label}
|
||||
self.is_presave = True
|
||||
|
||||
def to_log(self, prefix=None, image_postfix='jpg', video_postfix='mp4', rank = 0):
|
||||
if not self.is_presave:
|
||||
self.presave(prefix, image_postfix, video_postfix, rank = rank)
|
||||
if not self.is_presave:
|
||||
return self.data
|
||||
if isinstance(self.data, str):
|
||||
return self.data
|
||||
elif isinstance(self.data, dict):
|
||||
ret_data, ret_label = self.data["ret_data"], self.data["ret_label"]
|
||||
if self.build_html:
|
||||
html_prefix = prefix + '_probe.html'
|
||||
html_file = self.save_html(html_prefix, ret_data,
|
||||
ret_label)
|
||||
return {'ori_file': ret_data, 'html': html_file}
|
||||
else:
|
||||
return {'ori_file': ret_data}
|
||||
elif isinstance(self.data, list):
|
||||
ret_data, ret_label = [], []
|
||||
for one_data in self.data:
|
||||
if isinstance(one_data, dict):
|
||||
one_ret_data, one_ret_label = one_data["ret_data"], one_data["ret_label"]
|
||||
ret_data.extend(one_ret_data)
|
||||
ret_label.extend(one_ret_label)
|
||||
elif isinstance(one_data, str):
|
||||
ret_data.append(one_data)
|
||||
if (self.is_image or self.is_video) and self.build_html:
|
||||
html_prefix = prefix + '_probe.html'
|
||||
html_file = self.save_html(html_prefix, ret_data,
|
||||
ret_label)
|
||||
return {'ori_file': ret_data, 'html': html_file}
|
||||
else:
|
||||
return {'ori_file': ret_data}
|
||||
|
||||
@@ -0,0 +1,257 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class Media(Enum):
|
||||
TEXT = 1
|
||||
IMAGE = 2
|
||||
VIDEO = 3
|
||||
AUDIO = 4
|
||||
|
||||
|
||||
class HtmlVisualization(object):
|
||||
def __init__(
|
||||
self,
|
||||
allow_annotation=False,
|
||||
slice_size=1000,
|
||||
align='center',
|
||||
width_scale='60%',
|
||||
title='Visualization',
|
||||
):
|
||||
self.content_list = []
|
||||
self.rows_meta = []
|
||||
self.allow_annotation = allow_annotation
|
||||
self.slice_size = slice_size
|
||||
self.align = align
|
||||
self.width_scale = width_scale
|
||||
self.title = title
|
||||
self.html_start = '<html>'
|
||||
self.html_head = f'<head><meta charset="utf-8"><title>{title}</title></head>'
|
||||
self.html_style = '''
|
||||
<style>
|
||||
table {
|
||||
border-collapse: collapse;
|
||||
}
|
||||
td {
|
||||
width: "{width_scale}";
|
||||
align: "{align}";
|
||||
margin: 0px;
|
||||
border: 0px;
|
||||
padding: 0px;
|
||||
vertical-align: top;
|
||||
}
|
||||
video {
|
||||
margin: 0px;
|
||||
border: 0px solid #ccc;
|
||||
padding: 0px;
|
||||
}
|
||||
textarea {
|
||||
margin: 0px;
|
||||
border: 0px;
|
||||
padding: 0px;
|
||||
resize: none;
|
||||
border: 1px solid #ccc;
|
||||
}
|
||||
</style>
|
||||
<script>
|
||||
function adjustHeight() {
|
||||
const textareas = document.querySelectorAll('textarea');
|
||||
textareas.forEach(textarea => {
|
||||
const td = textarea.parentNode;
|
||||
const tdHeight = td.clientHeight;
|
||||
textarea.style.height = tdHeight + 'px';
|
||||
});
|
||||
}
|
||||
window.onload = adjustHeight;
|
||||
window.onresize = adjustHeight;
|
||||
</script>
|
||||
'''.replace('{width_scale}',
|
||||
self.width_scale).replace('{align}', self.align)
|
||||
self.html_body = '<body>{BODY}</body>\n'
|
||||
self.html_end = '</html>'
|
||||
self.html_script = '''
|
||||
<script>
|
||||
function saveSamples() {
|
||||
let selectedSamples = document.querySelectorAll('input[name="sample[]"]:checked');
|
||||
let notSelectedSamples = document.querySelectorAll('input[name="sample[]"]:not(:checked)');
|
||||
let sampleUrls = [];
|
||||
for (let i=0; i<selectedSamples.length; i++) {
|
||||
sampleUrls.push(selectedSamples[i].value + "#;#" + "1");
|
||||
}
|
||||
for (let i=0; i<notSelectedSamples.length; i++) {
|
||||
sampleUrls.push(notSelectedSamples[i].value + "#;#" + "0");
|
||||
}
|
||||
let fileContent = sampleUrls.join('\\n');
|
||||
let file = new Blob([fileContent], {type: 'text/plain'});
|
||||
let a = document.createElement('a');
|
||||
a.href = URL.createObjectURL(file);
|
||||
a.download = 'result.txt';
|
||||
a.click();
|
||||
}
|
||||
</script>
|
||||
'''
|
||||
self.label_button = (
|
||||
'<table><tr><td>' +
|
||||
"<button style='height: 50px;' type=\"button\" onclick=\"saveSamples()\">Save Samples</button>"
|
||||
+ '</td></tr></table>')
|
||||
|
||||
def format_col(self,
|
||||
content='',
|
||||
label='',
|
||||
type=Media.TEXT,
|
||||
content_height=400,
|
||||
content_width=600):
|
||||
if type == Media.TEXT:
|
||||
ret_str = '<td><textarea' # noqa: E501
|
||||
# if content_height is not None:
|
||||
# rows = f"rows={content_height//30}"
|
||||
# ret_str += f" {rows}"
|
||||
if content_width is not None:
|
||||
cols = f"cols={content_width//15}"
|
||||
ret_str += f" {cols}"
|
||||
ret_str += f'>"{content}"</textarea></td>\n'
|
||||
sec_ret_str = f'<td align="center"><font size="3"><strong>{label}<strong></font></td>\n'
|
||||
return [ret_str, sec_ret_str]
|
||||
elif type == Media.IMAGE:
|
||||
ret_str = f'<td><img src="{content}"'
|
||||
if content_height is not None:
|
||||
height = f'height="{content_height}"'
|
||||
ret_str += f" {height}"
|
||||
if content_width is not None:
|
||||
width = f'width="{content_width}"'
|
||||
ret_str += f" {width}"
|
||||
ret_str += ' ></td>\n'
|
||||
sec_ret_str = f'<td align="center"><font size="3"><strong>{label}<strong></font></td>\n'
|
||||
return [ret_str, sec_ret_str]
|
||||
elif type == Media.VIDEO:
|
||||
ret_str = '<td><video' # noqa
|
||||
if content_height is not None:
|
||||
height = f'height="{content_height}"'
|
||||
ret_str += f" {height}"
|
||||
if content_width is not None:
|
||||
width = f'width="{content_width}"'
|
||||
ret_str += f" {width}"
|
||||
ret_str += ' controls>'
|
||||
ret_str += f'<source src="{content}" type="video/mp4"></video></td>\n'
|
||||
sec_ret_str = f'<td align="center"><font size="3"><strong>{label}<strong></font></td>\n'
|
||||
return [ret_str, sec_ret_str]
|
||||
elif type == Media.AUDIO:
|
||||
ret_str = f'<td><audio src="{content}" controls></td>\n'
|
||||
sec_ret_str = f'<td align="center"><font size="3"><strong>{label}<strong></font></td>\n'
|
||||
return [ret_str, sec_ret_str]
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
def format_row(self):
|
||||
sample_id = 0
|
||||
all_sample_html = []
|
||||
for one_content, one_row_meta in zip(self.content_list,
|
||||
self.rows_meta):
|
||||
one_row_str = '<table><tr>'
|
||||
one_row_str += '\n'.join([v[0] for v in one_content])
|
||||
if self.allow_annotation:
|
||||
row_meta = '#;#'.join(one_row_meta)
|
||||
one_row_str += (
|
||||
f'<td><input type="checkbox" class="large-checkbox" '
|
||||
f'id="sample{sample_id}" name="sample[]" value="{row_meta}"></td>\n'
|
||||
)
|
||||
one_row_str += '</tr><tr>'
|
||||
one_row_str += '\n'.join([v[1] for v in one_content]) # noqa
|
||||
if self.allow_annotation: # noqa
|
||||
one_row_str += f'<td></td>\n' # noqa
|
||||
one_row_str += '</tr></table>'
|
||||
if self.allow_annotation:
|
||||
one_row_str = f'<label for="sample{sample_id}">{one_row_str}</label>'
|
||||
all_sample_html.append(one_row_str)
|
||||
sample_id += 1
|
||||
|
||||
return '\n'.join(all_sample_html)
|
||||
|
||||
def add_record(self,
|
||||
content='',
|
||||
label='',
|
||||
type=Media.TEXT,
|
||||
row_id=1,
|
||||
col_id=1,
|
||||
annotation_meta=None,
|
||||
content_height=None,
|
||||
content_width=None):
|
||||
if row_id >= len(self.content_list):
|
||||
self.content_list.append([])
|
||||
self.rows_meta.append([])
|
||||
if row_id != len(self.content_list) - 1:
|
||||
raise RuntimeError(
|
||||
'row_id should be next number of the last row_id.')
|
||||
if col_id > len(self.content_list[row_id]):
|
||||
raise RuntimeError(
|
||||
'col_id should be next number of the last col_id.')
|
||||
format_col = self.format_col(content, f"{row_id}-{col_id}: {label}",
|
||||
type, content_height, content_width)
|
||||
|
||||
annotation_meta = annotation_meta if annotation_meta else ''
|
||||
if col_id == len(self.content_list[row_id]):
|
||||
self.content_list[row_id].append(format_col)
|
||||
self.rows_meta[row_id].append(annotation_meta)
|
||||
else:
|
||||
self.content_list[row_id][col_id] = format_col
|
||||
self.rows_meta[row_id][col_id] = annotation_meta
|
||||
|
||||
def save_html(self, path):
|
||||
html_body = self.format_row()
|
||||
ret_html_list = [
|
||||
self.html_start, self.html_head, self.html_style,
|
||||
self.html_body.replace('{BODY}', html_body)
|
||||
]
|
||||
if self.allow_annotation:
|
||||
ret_html_list.append(self.label_button)
|
||||
ret_html_list.append(self.html_script)
|
||||
ret_html_list.append(self.html_end)
|
||||
ret_html = '\n'.join(ret_html_list)
|
||||
with open(path, 'w') as f:
|
||||
f.write(ret_html)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.file_system import FS
|
||||
FS.init_fs_client(Config(cfg_dict={}, load=False))
|
||||
|
||||
image_content_oss = '0_probe_0_[1024_2048_3].jpg'
|
||||
content_oss = '6UTWGRG1lx08iRBx5REA01041200dzcb0E010.mp4'
|
||||
caption = 'a little girl says hello.'
|
||||
|
||||
html_ins = HtmlVisualization(allow_annotation=True,
|
||||
slice_size=1000,
|
||||
title='Visualization',
|
||||
width_scale='100%')
|
||||
|
||||
for i in range(4):
|
||||
content_url = FS.get_url(content_oss, skip_check=True)
|
||||
html_ins.add_record(content=content_url,
|
||||
label='caption',
|
||||
type=Media.VIDEO,
|
||||
row_id=i,
|
||||
col_id=0,
|
||||
annotation_meta=None,
|
||||
content_height=600,
|
||||
content_width=None)
|
||||
html_ins.add_record(content=caption,
|
||||
label='caption',
|
||||
type=Media.TEXT,
|
||||
row_id=i,
|
||||
col_id=1,
|
||||
annotation_meta=None,
|
||||
content_height=600,
|
||||
content_width=750)
|
||||
image_content_url = FS.get_url(image_content_oss, skip_check=True)
|
||||
html_ins.add_record(content=image_content_url,
|
||||
label='caption',
|
||||
type=Media.IMAGE,
|
||||
row_id=i,
|
||||
col_id=2,
|
||||
annotation_meta=None,
|
||||
content_height=600,
|
||||
content_width=None)
|
||||
|
||||
with FS.put_to('visualize.html') as local_path:
|
||||
html_ins.save_html(local_path)
|
||||
Reference in New Issue
Block a user