update v1.1.0

This commit is contained in:
zeyinzi.jzyz
2024-10-21 00:35:53 +08:00
parent 7d6451efad
commit 0bba2c319d
148 changed files with 18476 additions and 1356 deletions
+12
View File
@@ -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
+135
View File
@@ -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)
+51
View File
@@ -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)
+38
View File
@@ -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
+271
View File
@@ -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)
+98
View File
@@ -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)
+3 -2
View File
@@ -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
+189
View File
@@ -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)
+935
View File
@@ -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)
+382
View File
@@ -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)
+160
View File
@@ -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
+3 -1
View File
@@ -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)
+5 -2
View File
@@ -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:
+10 -6
View File
@@ -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
+216
View File
@@ -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)
+11 -5
View File
@@ -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 \
+1 -1
View File
@@ -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 -1
View File
@@ -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
+245
View File
@@ -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_)
+7 -1
View File
@@ -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)
+211
View File
@@ -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)
+10
View File
@@ -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)
+2 -1
View File
@@ -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)
+16 -8
View File
@@ -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,
+27
View File
@@ -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,
+59 -20
View File
@@ -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)
+355 -175
View File
@@ -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}.')
+75 -6
View File
@@ -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()
+33 -24
View File
@@ -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:
+70 -29
View File
@@ -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
+18 -8
View File
@@ -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)
+32 -15
View File
@@ -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):
+7 -1
View File
@@ -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
+362 -15
View File
@@ -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]:
+19 -6
View File
@@ -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 == '':
+7 -4
View File
@@ -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:
+51 -1
View File
@@ -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)
+3 -3
View File
@@ -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
View File
@@ -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('<', '&lt;').replace(
'>', '&gt;')
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}
+257
View File
@@ -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)