From 037e73fe6e67272bb6758be85e9b632ad03a1028 Mon Sep 17 00:00:00 2001 From: kangzhao2 Date: Tue, 15 Aug 2023 21:32:30 +0800 Subject: [PATCH] baishao --- modelscope/metainfo.py | 4 +- .../multi_modal/image_to_video/__init__.py | 23 + .../image_to_video/image_to_video_model.py | 151 ++ .../image_to_video/modules/__init__.py | 3 + .../image_to_video/modules/autoencoder.py | 621 +++++++++ .../image_to_video/modules/embedder.py | 77 + .../image_to_video/modules/unet_i2v.py | 1236 +++++++++++++++++ .../image_to_video/utils/__init__.py | 1 + .../image_to_video/utils/config.py | 166 +++ .../image_to_video/utils/diffusion.py | 377 +++++ .../image_to_video/utils/registry.py | 155 +++ .../utils/registry_class/__init__.py | 4 + .../utils/registry_class/autoencoder.py | 11 + .../utils/registry_class/distrubution.py | 11 + .../utils/registry_class/embedder.py | 11 + .../utils/registry_class/model.py | 11 + .../multi_modal/image_to_video/utils/seed.py | 12 + .../image_to_video/utils/shedule.py | 38 + .../image_to_video/utils/transforms.py | 377 +++++ .../multi_modal/image_to_video_pipeline.py | 74 + modelscope/utils/constant.py | 1 + requirements.txt | 1 - test_image2video.py | 38 + tests/pipelines/test_image2video.py | 38 + 24 files changed, 3438 insertions(+), 3 deletions(-) create mode 100644 modelscope/models/multi_modal/image_to_video/__init__.py create mode 100755 modelscope/models/multi_modal/image_to_video/image_to_video_model.py create mode 100755 modelscope/models/multi_modal/image_to_video/modules/__init__.py create mode 100755 modelscope/models/multi_modal/image_to_video/modules/autoencoder.py create mode 100755 modelscope/models/multi_modal/image_to_video/modules/embedder.py create mode 100755 modelscope/models/multi_modal/image_to_video/modules/unet_i2v.py create mode 100755 modelscope/models/multi_modal/image_to_video/utils/__init__.py create mode 100755 modelscope/models/multi_modal/image_to_video/utils/config.py create mode 100755 modelscope/models/multi_modal/image_to_video/utils/diffusion.py create mode 100755 modelscope/models/multi_modal/image_to_video/utils/registry.py create mode 100755 modelscope/models/multi_modal/image_to_video/utils/registry_class/__init__.py create mode 100755 modelscope/models/multi_modal/image_to_video/utils/registry_class/autoencoder.py create mode 100755 modelscope/models/multi_modal/image_to_video/utils/registry_class/distrubution.py create mode 100755 modelscope/models/multi_modal/image_to_video/utils/registry_class/embedder.py create mode 100755 modelscope/models/multi_modal/image_to_video/utils/registry_class/model.py create mode 100755 modelscope/models/multi_modal/image_to_video/utils/seed.py create mode 100644 modelscope/models/multi_modal/image_to_video/utils/shedule.py create mode 100755 modelscope/models/multi_modal/image_to_video/utils/transforms.py create mode 100644 modelscope/pipelines/multi_modal/image_to_video_pipeline.py create mode 100644 test_image2video.py create mode 100644 tests/pipelines/test_image2video.py diff --git a/modelscope/metainfo.py b/modelscope/metainfo.py index 70fd9c86..eb6a801c 100644 --- a/modelscope/metainfo.py +++ b/modelscope/metainfo.py @@ -218,8 +218,8 @@ class Models(object): mplug_owl = 'mplug-owl' clip_interrogator = 'clip-interrogator' stable_diffusion = 'stable-diffusion' - videocomposer = 'videocomposer' text_to_360panorama_image = 'text-to-360panorama-image' + image_to_video_model = 'image-to-video-model' # science models unifold = 'unifold' @@ -526,7 +526,6 @@ class Pipelines(object): multi_modal_similarity = 'multi-modal-similarity' text_to_image_synthesis = 'text-to-image-synthesis' video_multi_modal_embedding = 'video-multi-modal-embedding' - videocomposer = 'videocomposer' image_text_retrieval = 'image-text-retrieval' ofa_ocr_recognition = 'ofa-ocr-recognition' ofa_asr = 'ofa-asr' @@ -545,6 +544,7 @@ class Pipelines(object): efficient_diffusion_tuning = 'efficient-diffusion-tuning' multimodal_dialogue = 'multimodal-dialogue' llama2_text_generation_pipeline = 'llama2-text-generation-pipeline' + image_to_video_task_pipeline = 'image-to-video-task-pipeline' # science tasks protein_structure = 'unifold-protein-structure' diff --git a/modelscope/models/multi_modal/image_to_video/__init__.py b/modelscope/models/multi_modal/image_to_video/__init__.py new file mode 100644 index 00000000..88155033 --- /dev/null +++ b/modelscope/models/multi_modal/image_to_video/__init__.py @@ -0,0 +1,23 @@ +# Copyright (c) Alibaba, Inc. and its affiliates. +from typing import TYPE_CHECKING + +from modelscope.utils.import_utils import LazyImportModule + +if TYPE_CHECKING: + + from .image_to_video_model import ImageToVideo + +else: + _import_structure = { + 'image_to_video_model': ['ImageToVideo'], + } + + import sys + + sys.modules[__name__] = LazyImportModule( + __name__, + globals()['__file__'], + _import_structure, + module_spec=__spec__, + extra_objects={}, + ) diff --git a/modelscope/models/multi_modal/image_to_video/image_to_video_model.py b/modelscope/models/multi_modal/image_to_video/image_to_video_model.py new file mode 100755 index 00000000..a388718b --- /dev/null +++ b/modelscope/models/multi_modal/image_to_video/image_to_video_model.py @@ -0,0 +1,151 @@ +import os +import os.path as osp +import random +import torch +from PIL import Image +import torch.cuda.amp as amp +from copy import copy +from typing import Any, Dict + +from modelscope.utils.constant import ModelFile, Tasks +from modelscope.metainfo import Models +from modelscope.models.base import TorchModel +from modelscope.models.builder import MODELS +from modelscope.utils.config import Config +from modelscope.utils.logger import get_logger + +from modelscope.models.multi_modal.image_to_video.modules import * +from modelscope.models.multi_modal.image_to_video.utils.config import cfg +from modelscope.models.multi_modal.image_to_video.utils.diffusion import GaussianDiffusion +import modelscope.models.multi_modal.image_to_video.utils.transforms as data +from modelscope.models.multi_modal.image_to_video.utils.seed import setup_seed +from modelscope.models.multi_modal.image_to_video.utils.shedule import beta_schedule +from modelscope.models.multi_modal.image_to_video.utils.registry_class import UNET, EMBEDDER, AUTO_ENCODER + +__all__ = ['ImageToVideo'] + +logger = get_logger() + +@MODELS.register_module(Tasks.image_to_video_task, module_name=Models.image_to_video_model) +class ImageToVideo(TorchModel): + def __init__(self, model_dir, *args, **kwargs): + super().__init__(model_dir=model_dir, *args, **kwargs) + + self.config = Config.from_file(osp.join(model_dir, ModelFile.CONFIGURATION)) + + # assign default value + cfg.batch_size = 1 + cfg.target_fps = 8 + cfg.max_frames = 32 + cfg.latent_hei = 32 + cfg.latent_wid = 56 + cfg.model_path = osp.join(model_dir, self.config.model.model_args.ckpt_unet) + + self.device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu') + + if 'seed' in self.config.model.model_args.keys(): + cfg.seed = self.config.model.model_args.seed + else: + cfg.seed = random.randint(0, 99999) + setup_seed(cfg.seed) + + # transform + vid_trans = data.Compose([ + data.CenterCropWide(size=(cfg.resolution[0], cfg.resolution[0])), + data.Resize(cfg.vit_resolution), + data.ToTensor(), + data.Normalize(mean=cfg.vit_mean, std=cfg.vit_std)]) + self.vid_trans = vid_trans + + cfg.embedder.pretrained = osp.join(model_dir, self.config.model.model_args.ckpt_clip) + clip_encoder = EMBEDDER.build(cfg.embedder) + clip_encoder.model.to(self.device) + self.clip_encoder = clip_encoder + logger.info(f'Build encoder with {cfg.embedder.type}') + + # [unet] + generator = UNET.build(cfg.UNet) + generator = generator.to(self.device) + generator.eval() + load_dict = torch.load(cfg.model_path, map_location='cpu') + ret = generator.load_state_dict(load_dict['state_dict'], strict=True) + self.generator = generator + logger.info('Load model {} path {}, with local status {}'.format(cfg.UNet.type, cfg.model_path, ret)) + + # [diffusion] + betas = beta_schedule('linear_sd', cfg.num_timesteps, init_beta=0.00085, last_beta=0.0120) + diffusion = GaussianDiffusion( + betas=betas, + mean_type=cfg.mean_type, + var_type=cfg.var_type, + loss_type=cfg.loss_type, + rescale_timesteps=False, + noise_strength=getattr(cfg, 'noise_strength', 0)) + self.diffusion = diffusion + logger.info('Build diffusion with type of GaussianDiffusion') + + # [auotoencoder] + cfg.auto_encoder.pretrained = osp.join(model_dir, self.config.model.model_args.ckpt_autoencoder) + autoencoder = AUTO_ENCODER.build(cfg.auto_encoder) + autoencoder.eval() # freeze + for param in autoencoder.parameters(): + param.requires_grad = False + autoencoder.to(self.device) + self.autoencoder = autoencoder + torch.cuda.empty_cache() + + zero_feature = torch.zeros(1, 1, cfg.UNet.input_dim).to(self.device) + self.zero_feature = zero_feature + self.fps_tensor = torch.tensor([cfg.target_fps], dtype=torch.long, device=self.device) + self.cfg = cfg + + def forward(self, input: Dict[str, Any]): + img_path = input['img_path'] + + cfg = self.cfg + image = Image.open(img_path) + if image.mode != 'RGB': + image = image.convert('RGB') + + vit_frame = self.vid_trans(image) + vit_frame = vit_frame.unsqueeze(0) + vit_frame = vit_frame.to(self.device) + img_embedding = self.clip_encoder(vit_frame).unsqueeze(1) + + noise = self.build_noise() + zero_feature = copy(self.zero_feature) + with torch.no_grad(): + with amp.autocast(enabled=cfg.use_fp16): + model_kwargs=[ + {'y': img_embedding, 'fps': self.fps_tensor}, + {'y': zero_feature.repeat(cfg.batch_size, 1, 1), 'fps': self.fps_tensor}] + gen_video = self.diffusion.ddim_sample_loop( + noise=noise, + model=self.generator, + model_kwargs=model_kwargs, + guide_scale=cfg.guide_scale, + ddim_timesteps=cfg.ddim_timesteps, + eta=0.0) + + gen_video = 1. / cfg.scale_factor * gen_video # [1, 4, 32, 32, 56] + gen_video = rearrange(gen_video, 'b c f h w -> (b f) c h w') + chunk_size = min(cfg.decoder_bs, gen_video.shape[0]) + gen_video_list = torch.chunk(gen_video, gen_video.shape[0]//chunk_size, dim=0) + decode_generator = [] + for vd_data in gen_video_list: + gen_frames = self.autoencoder.decode(vd_data) + decode_generator.append(gen_frames) + + gen_video = torch.cat(decode_generator, dim=0) + gen_video = rearrange(gen_video, '(b f) c h w -> b c f h w', b = cfg.batch_size) + + return gen_video.type(torch.float32).cpu() + + def build_noise(self): + cfg = self.cfg + noise = torch.randn([1, 4, cfg.max_frames, cfg.latent_hei, cfg.latent_wid]).to(self.device) + if cfg.noise_strength > 0: + b, c, f, *_ = noise.shape + offset_noise = torch.randn(b, c, f, 1, 1, device=noise.device) + noise = noise + cfg.noise_strength * offset_noise + return noise.contiguous() diff --git a/modelscope/models/multi_modal/image_to_video/modules/__init__.py b/modelscope/models/multi_modal/image_to_video/modules/__init__.py new file mode 100755 index 00000000..2bfd7e28 --- /dev/null +++ b/modelscope/models/multi_modal/image_to_video/modules/__init__.py @@ -0,0 +1,3 @@ +from .embedder import * +from .unet_i2v import * +from .autoencoder import * \ No newline at end of file diff --git a/modelscope/models/multi_modal/image_to_video/modules/autoencoder.py b/modelscope/models/multi_modal/image_to_video/modules/autoencoder.py new file mode 100755 index 00000000..e482f680 --- /dev/null +++ b/modelscope/models/multi_modal/image_to_video/modules/autoencoder.py @@ -0,0 +1,621 @@ +import torch +import logging +import collections +import numpy as np +import torch.nn as nn +import torch.nn.functional as F + +from ..utils.registry_class import AUTO_ENCODER +from ..utils.registry_class import DISTRIBUTION + +def nonlinearity(x): + # swish + return x*torch.sigmoid(x) + +def Normalize(in_channels, num_groups=32): + return torch.nn.GroupNorm(num_groups=num_groups, num_channels=in_channels, eps=1e-6, affine=True) + + +@DISTRIBUTION.register_class() +class DiagonalGaussianDistribution(object): + def __init__(self, parameters, deterministic=False): + self.parameters = parameters + self.mean, self.logvar = torch.chunk(parameters, 2, dim=1) + self.logvar = torch.clamp(self.logvar, -30.0, 20.0) + self.deterministic = deterministic + self.std = torch.exp(0.5 * self.logvar) + self.var = torch.exp(self.logvar) + if self.deterministic: + self.var = self.std = torch.zeros_like(self.mean).to(device=self.parameters.device) + + def sample(self): + x = self.mean + self.std * torch.randn(self.mean.shape).to(device=self.parameters.device) + return x + + def kl(self, other=None): + if self.deterministic: + return torch.Tensor([0.]) + else: + if other is None: + return 0.5 * torch.sum(torch.pow(self.mean, 2) + + self.var - 1.0 - self.logvar, + dim=[1, 2, 3]) + else: + return 0.5 * torch.sum( + torch.pow(self.mean - other.mean, 2) / other.var + + self.var / other.var - 1.0 - self.logvar + other.logvar, + dim=[1, 2, 3]) + + def nll(self, sample, dims=[1,2,3]): + if self.deterministic: + return torch.Tensor([0.]) + logtwopi = np.log(2.0 * np.pi) + return 0.5 * torch.sum( + logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var, + dim=dims) + + def mode(self): + return self.mean + +class Downsample(nn.Module): + def __init__(self, in_channels, with_conv): + super().__init__() + self.with_conv = with_conv + if self.with_conv: + # no asymmetric padding in torch conv, must do it ourselves + self.conv = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=3, + stride=2, + padding=0) + + def forward(self, x): + if self.with_conv: + pad = (0,1,0,1) + x = torch.nn.functional.pad(x, pad, mode="constant", value=0) + x = self.conv(x) + else: + x = torch.nn.functional.avg_pool2d(x, kernel_size=2, stride=2) + return x + +class ResnetBlock(nn.Module): + def __init__(self, *, in_channels, out_channels=None, conv_shortcut=False, + dropout, temb_channels=512): + super().__init__() + self.in_channels = in_channels + out_channels = in_channels if out_channels is None else out_channels + self.out_channels = out_channels + self.use_conv_shortcut = conv_shortcut + + self.norm1 = Normalize(in_channels) + self.conv1 = torch.nn.Conv2d(in_channels, + out_channels, + kernel_size=3, + stride=1, + padding=1) + if temb_channels > 0: + self.temb_proj = torch.nn.Linear(temb_channels, + out_channels) + self.norm2 = Normalize(out_channels) + self.dropout = torch.nn.Dropout(dropout) + self.conv2 = torch.nn.Conv2d(out_channels, + out_channels, + kernel_size=3, + stride=1, + padding=1) + if self.in_channels != self.out_channels: + if self.use_conv_shortcut: + self.conv_shortcut = torch.nn.Conv2d(in_channels, + out_channels, + kernel_size=3, + stride=1, + padding=1) + else: + self.nin_shortcut = torch.nn.Conv2d(in_channels, + out_channels, + kernel_size=1, + stride=1, + padding=0) + + def forward(self, x, temb): + h = x + h = self.norm1(h) + h = nonlinearity(h) + h = self.conv1(h) + + if temb is not None: + h = h + self.temb_proj(nonlinearity(temb))[:,:,None,None] + + h = self.norm2(h) + h = nonlinearity(h) + h = self.dropout(h) + h = self.conv2(h) + + if self.in_channels != self.out_channels: + if self.use_conv_shortcut: + x = self.conv_shortcut(x) + else: + x = self.nin_shortcut(x) + + return x+h + + +class AttnBlock(nn.Module): + def __init__(self, in_channels): + super().__init__() + self.in_channels = in_channels + + self.norm = Normalize(in_channels) + self.q = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=1, + stride=1, + padding=0) + self.k = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=1, + stride=1, + padding=0) + self.v = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=1, + stride=1, + padding=0) + self.proj_out = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=1, + stride=1, + padding=0) + + def forward(self, x): + h_ = x + h_ = self.norm(h_) + q = self.q(h_) + k = self.k(h_) + v = self.v(h_) + + # compute attention + b,c,h,w = q.shape + q = q.reshape(b,c,h*w) + q = q.permute(0,2,1) # b,hw,c + k = k.reshape(b,c,h*w) # b,c,hw + w_ = torch.bmm(q,k) # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j] + w_ = w_ * (int(c)**(-0.5)) + w_ = torch.nn.functional.softmax(w_, dim=2) + + # attend to values + v = v.reshape(b,c,h*w) + w_ = w_.permute(0,2,1) # b,hw,hw (first hw of k, second of q) + h_ = torch.bmm(v,w_) # b, c,hw (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j] + h_ = h_.reshape(b,c,h,w) + + h_ = self.proj_out(h_) + + return x+h_ + +class AttnBlock(nn.Module): + def __init__(self, in_channels): + super().__init__() + self.in_channels = in_channels + + self.norm = Normalize(in_channels) + self.q = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=1, + stride=1, + padding=0) + self.k = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=1, + stride=1, + padding=0) + self.v = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=1, + stride=1, + padding=0) + self.proj_out = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=1, + stride=1, + padding=0) + + def forward(self, x): + h_ = x + h_ = self.norm(h_) + q = self.q(h_) + k = self.k(h_) + v = self.v(h_) + + # compute attention + b,c,h,w = q.shape + q = q.reshape(b,c,h*w) + q = q.permute(0,2,1) # b,hw,c + k = k.reshape(b,c,h*w) # b,c,hw + w_ = torch.bmm(q,k) # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j] + w_ = w_ * (int(c)**(-0.5)) + w_ = torch.nn.functional.softmax(w_, dim=2) + + # attend to values + v = v.reshape(b,c,h*w) + w_ = w_.permute(0,2,1) # b,hw,hw (first hw of k, second of q) + h_ = torch.bmm(v,w_) # b, c,hw (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j] + h_ = h_.reshape(b,c,h,w) + + h_ = self.proj_out(h_) + + return x+h_ + +class Upsample(nn.Module): + def __init__(self, in_channels, with_conv): + super().__init__() + self.with_conv = with_conv + if self.with_conv: + self.conv = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=3, + stride=1, + padding=1) + + def forward(self, x): + x = torch.nn.functional.interpolate(x, scale_factor=2.0, mode="nearest") + if self.with_conv: + x = self.conv(x) + return x + + +class Downsample(nn.Module): + def __init__(self, in_channels, with_conv): + super().__init__() + self.with_conv = with_conv + if self.with_conv: + # no asymmetric padding in torch conv, must do it ourselves + self.conv = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=3, + stride=2, + padding=0) + + def forward(self, x): + if self.with_conv: + pad = (0,1,0,1) + x = torch.nn.functional.pad(x, pad, mode="constant", value=0) + x = self.conv(x) + else: + x = torch.nn.functional.avg_pool2d(x, kernel_size=2, stride=2) + return x + +class Encoder(nn.Module): + def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks, + attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels, + resolution, z_channels, double_z=True, use_linear_attn=False, attn_type="vanilla", + **ignore_kwargs): + super().__init__() + if use_linear_attn: attn_type = "linear" + self.ch = ch + self.temb_ch = 0 + self.num_resolutions = len(ch_mult) + self.num_res_blocks = num_res_blocks + self.resolution = resolution + self.in_channels = in_channels + + # downsampling + self.conv_in = torch.nn.Conv2d(in_channels, + self.ch, + kernel_size=3, + stride=1, + padding=1) + + curr_res = resolution + in_ch_mult = (1,)+tuple(ch_mult) + self.in_ch_mult = in_ch_mult + self.down = nn.ModuleList() + for i_level in range(self.num_resolutions): + block = nn.ModuleList() + attn = nn.ModuleList() + block_in = ch*in_ch_mult[i_level] + block_out = ch*ch_mult[i_level] + for i_block in range(self.num_res_blocks): + block.append(ResnetBlock(in_channels=block_in, + out_channels=block_out, + temb_channels=self.temb_ch, + dropout=dropout)) + block_in = block_out + if curr_res in attn_resolutions: + attn.append(AttnBlock(block_in)) + down = nn.Module() + down.block = block + down.attn = attn + if i_level != self.num_resolutions-1: + down.downsample = Downsample(block_in, resamp_with_conv) + curr_res = curr_res // 2 + self.down.append(down) + + # middle + self.mid = nn.Module() + self.mid.block_1 = ResnetBlock(in_channels=block_in, + out_channels=block_in, + temb_channels=self.temb_ch, + dropout=dropout) + self.mid.attn_1 = AttnBlock(block_in) + self.mid.block_2 = ResnetBlock(in_channels=block_in, + out_channels=block_in, + temb_channels=self.temb_ch, + dropout=dropout) + + # end + self.norm_out = Normalize(block_in) + self.conv_out = torch.nn.Conv2d(block_in, + 2*z_channels if double_z else z_channels, + kernel_size=3, + stride=1, + padding=1) + + def forward(self, x): + # timestep embedding + temb = None + + # downsampling + hs = [self.conv_in(x)] + for i_level in range(self.num_resolutions): + for i_block in range(self.num_res_blocks): + h = self.down[i_level].block[i_block](hs[-1], temb) + if len(self.down[i_level].attn) > 0: + h = self.down[i_level].attn[i_block](h) + hs.append(h) + if i_level != self.num_resolutions-1: + hs.append(self.down[i_level].downsample(hs[-1])) + + # middle + h = hs[-1] + h = self.mid.block_1(h, temb) + h = self.mid.attn_1(h) + h = self.mid.block_2(h, temb) + + # end + h = self.norm_out(h) + h = nonlinearity(h) + h = self.conv_out(h) + return h + + +class Decoder(nn.Module): + def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks, + attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels, + resolution, z_channels, give_pre_end=False, tanh_out=False, use_linear_attn=False, + attn_type="vanilla", **ignorekwargs): + super().__init__() + if use_linear_attn: attn_type = "linear" + self.ch = ch + self.temb_ch = 0 + self.num_resolutions = len(ch_mult) + self.num_res_blocks = num_res_blocks + self.resolution = resolution + self.in_channels = in_channels + self.give_pre_end = give_pre_end + self.tanh_out = tanh_out + + # compute in_ch_mult, block_in and curr_res at lowest res + in_ch_mult = (1,)+tuple(ch_mult) + block_in = ch*ch_mult[self.num_resolutions-1] + curr_res = resolution // 2**(self.num_resolutions-1) + self.z_shape = (1,z_channels, curr_res, curr_res) + logging.info("Working with z of shape {} = {} dimensions.".format( + self.z_shape, np.prod(self.z_shape))) + + # z to block_in + self.conv_in = torch.nn.Conv2d(z_channels, + block_in, + kernel_size=3, + stride=1, + padding=1) + + # middle + self.mid = nn.Module() + self.mid.block_1 = ResnetBlock(in_channels=block_in, + out_channels=block_in, + temb_channels=self.temb_ch, + dropout=dropout) + self.mid.attn_1 = AttnBlock(block_in) + self.mid.block_2 = ResnetBlock(in_channels=block_in, + out_channels=block_in, + temb_channels=self.temb_ch, + dropout=dropout) + + # upsampling + self.up = nn.ModuleList() + for i_level in reversed(range(self.num_resolutions)): + block = nn.ModuleList() + attn = nn.ModuleList() + block_out = ch*ch_mult[i_level] + for i_block in range(self.num_res_blocks+1): + block.append(ResnetBlock(in_channels=block_in, + out_channels=block_out, + temb_channels=self.temb_ch, + dropout=dropout)) + block_in = block_out + if curr_res in attn_resolutions: + attn.append(AttnBlock(block_in)) + up = nn.Module() + up.block = block + up.attn = attn + if i_level != 0: + up.upsample = Upsample(block_in, resamp_with_conv) + curr_res = curr_res * 2 + self.up.insert(0, up) # prepend to get consistent order + + # end + self.norm_out = Normalize(block_in) + self.conv_out = torch.nn.Conv2d(block_in, + out_ch, + kernel_size=3, + stride=1, + padding=1) + + def forward(self, z): + #assert z.shape[1:] == self.z_shape[1:] + self.last_z_shape = z.shape + + # timestep embedding + temb = None + + # z to block_in + h = self.conv_in(z) + + # middle + h = self.mid.block_1(h, temb) + h = self.mid.attn_1(h) + h = self.mid.block_2(h, temb) + + # upsampling + for i_level in reversed(range(self.num_resolutions)): + for i_block in range(self.num_res_blocks+1): + h = self.up[i_level].block[i_block](h, temb) + if len(self.up[i_level].attn) > 0: + h = self.up[i_level].attn[i_block](h) + if i_level != 0: + h = self.up[i_level].upsample(h) + + # end + if self.give_pre_end: + return h + + h = self.norm_out(h) + h = nonlinearity(h) + h = self.conv_out(h) + if self.tanh_out: + h = torch.tanh(h) + return h + +@AUTO_ENCODER.register_class() +class AutoencoderKL(nn.Module): + def __init__(self, + ddconfig, + embed_dim, + pretrained=None, + ignore_keys=[], + image_key="image", + colorize_nlabels=None, + monitor=None, + ema_decay=None, + learn_logvar=False, + **kwargs): + super().__init__() + self.learn_logvar = learn_logvar + self.image_key = image_key + self.encoder = Encoder(**ddconfig) + self.decoder = Decoder(**ddconfig) + assert ddconfig["double_z"] + self.quant_conv = torch.nn.Conv2d(2*ddconfig["z_channels"], 2*embed_dim, 1) + self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig["z_channels"], 1) + self.embed_dim = embed_dim + if colorize_nlabels is not None: + assert type(colorize_nlabels)==int + self.register_buffer("colorize", torch.randn(3, colorize_nlabels, 1, 1)) + if monitor is not None: + self.monitor = monitor + + self.use_ema = ema_decay is not None + + if pretrained is not None: + self.init_from_ckpt(pretrained, ignore_keys=ignore_keys) + + def init_from_ckpt(self, path, ignore_keys=list()): + sd = torch.load(path, map_location="cpu")["state_dict"] + keys = list(sd.keys()) + sd_new = collections.OrderedDict() + for k in keys: + if k.find('first_stage_model') >= 0: + k_new = k.split('first_stage_model.')[-1] + sd_new[k_new] = sd[k] + self.load_state_dict(sd_new, strict=True) + logging.info(f"Restored from {path}") + + def on_train_batch_end(self, *args, **kwargs): + if self.use_ema: + self.model_ema(self) + + def encode(self, x): + h = self.encoder(x) + moments = self.quant_conv(h) + posterior = DiagonalGaussianDistribution(moments) + return posterior + + def decode(self, z): + z = self.post_quant_conv(z) + dec = self.decoder(z) + return dec + + def forward(self, input, sample_posterior=True): + posterior = self.encode(input) + if sample_posterior: + z = posterior.sample() + else: + z = posterior.mode() + dec = self.decode(z) + return dec, posterior + + def get_input(self, batch, k): + x = batch[k] + if len(x.shape) == 3: + x = x[..., None] + x = x.permute(0, 3, 1, 2).to(memory_format=torch.contiguous_format).float() + return x + + def get_last_layer(self): + return self.decoder.conv_out.weight + + @torch.no_grad() + def log_images(self, batch, only_inputs=False, log_ema=False, **kwargs): + log = dict() + x = self.get_input(batch, self.image_key) + x = x.to(self.device) + if not only_inputs: + xrec, posterior = self(x) + if x.shape[1] > 3: + # colorize with random projection + assert xrec.shape[1] > 3 + x = self.to_rgb(x) + xrec = self.to_rgb(xrec) + log["samples"] = self.decode(torch.randn_like(posterior.sample())) + log["reconstructions"] = xrec + if log_ema or self.use_ema: + with self.ema_scope(): + xrec_ema, posterior_ema = self(x) + if x.shape[1] > 3: + # colorize with random projection + assert xrec_ema.shape[1] > 3 + xrec_ema = self.to_rgb(xrec_ema) + log["samples_ema"] = self.decode(torch.randn_like(posterior_ema.sample())) + log["reconstructions_ema"] = xrec_ema + log["inputs"] = x + return log + + def to_rgb(self, x): + assert self.image_key == "segmentation" + if not hasattr(self, "colorize"): + self.register_buffer("colorize", torch.randn(3, x.shape[1], 1, 1).to(x)) + x = F.conv2d(x, weight=self.colorize) + x = 2.*(x-x.min())/(x.max()-x.min()) - 1. + return x + + +class IdentityFirstStage(torch.nn.Module): + def __init__(self, *args, vq_interface=False, **kwargs): + self.vq_interface = vq_interface + super().__init__() + + def encode(self, x, *args, **kwargs): + return x + + def decode(self, x, *args, **kwargs): + return x + + def quantize(self, x, *args, **kwargs): + if self.vq_interface: + return x, None, [None, None, None] + return x + + def forward(self, x, *args, **kwargs): + return x + diff --git a/modelscope/models/multi_modal/image_to_video/modules/embedder.py b/modelscope/models/multi_modal/image_to_video/modules/embedder.py new file mode 100755 index 00000000..f974b59f --- /dev/null +++ b/modelscope/models/multi_modal/image_to_video/modules/embedder.py @@ -0,0 +1,77 @@ +import os +import torch +import logging +import open_clip +import numpy as np +import torch.nn as nn +import torchvision.transforms as T + +from ..utils.registry_class import EMBEDDER + +@EMBEDDER.register_class() +class FrozenOpenCLIPVisualEmbedder(nn.Module): + """ + Uses the OpenCLIP transformer encoder for text + """ + LAYERS = [ + #"pooled", + "last", + "penultimate" + ] + def __init__(self, pretrained, vit_resolution=(224, 224), arch="ViT-H-14", device="cuda", max_length=77, + freeze=True, layer="last"): + super().__init__() + assert layer in self.LAYERS + model, _, preprocess = open_clip.create_model_and_transforms( + arch, device=torch.device('cpu'), pretrained=pretrained) + # Normalize(mean=(0.48145466, 0.4578275, 0.40821073), std=(0.26862954, 0.26130258, 0.27577711) + del model.transformer + self.model = model + data_white = np.ones((vit_resolution[0], vit_resolution[1], 3), dtype=np.uint8)*255 + self.white_image = preprocess(T.ToPILImage()(data_white)).unsqueeze(0) + + self.device = device + self.max_length = max_length # 77 + if freeze: + self.freeze() + self.layer = layer # 'penultimate' + if self.layer == "last": + self.layer_idx = 0 + elif self.layer == "penultimate": + self.layer_idx = 1 + else: + raise NotImplementedError() + + def freeze(self): # model.encode_image(torch.randn(2,3,224,224)) + self.model = self.model.eval() + for param in self.parameters(): + param.requires_grad = False + + def forward(self, image): + # tokens = open_clip.tokenize(text) + z = self.model.encode_image(image.to(self.device)) + return z + + def encode_with_transformer(self, text): + x = self.model.token_embedding(text) # [batch_size, n_ctx, d_model] + x = x + self.model.positional_embedding + x = x.permute(1, 0, 2) # NLD -> LND + x = self.text_transformer_forward(x, attn_mask=self.model.attn_mask) + x = x.permute(1, 0, 2) # LND -> NLD + x = self.model.ln_final(x) + + return x + + def text_transformer_forward(self, x: torch.Tensor, attn_mask = None): + for i, r in enumerate(self.model.transformer.resblocks): + if i == len(self.model.transformer.resblocks) - self.layer_idx: + break + if self.model.transformer.grad_checkpointing and not torch.jit.is_scripting(): + x = checkpoint(r, x, attn_mask) + else: + x = r(x, attn_mask=attn_mask) + return x + + def encode(self, text): + return self(text) + diff --git a/modelscope/models/multi_modal/image_to_video/modules/unet_i2v.py b/modelscope/models/multi_modal/image_to_video/modules/unet_i2v.py new file mode 100755 index 00000000..6a9dee96 --- /dev/null +++ b/modelscope/models/multi_modal/image_to_video/modules/unet_i2v.py @@ -0,0 +1,1236 @@ +import math +import torch +import torch.nn as nn +from einops import rearrange +import torch.nn.functional as F +from rotary_embedding_torch import RotaryEmbedding +from fairscale.nn.checkpoint import checkpoint_wrapper + +from ..utils.registry_class import UNET + +USE_TEMPORAL_TRANSFORMER = True + +def sinusoidal_embedding(timesteps, dim): + # check input + half = dim // 2 + timesteps = timesteps.float() + + # compute sinusoidal embedding + sinusoid = torch.outer( + timesteps, + torch.pow(10000, -torch.arange(half).to(timesteps).div(half))) + x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1) + if dim % 2 != 0: + x = torch.cat([x, torch.zeros_like(x[:, :1])], dim=1) + return x + +def exists(x): + return x is not None + +def default(val, d): + if exists(val): + return val + return d() if callable(d) else d + + +def prob_mask_like(shape, prob, device): + if prob == 1: + return torch.ones(shape, device = device, dtype = torch.bool) + elif prob == 0: + return torch.zeros(shape, device = device, dtype = torch.bool) + else: + mask = torch.zeros(shape, device = device).float().uniform_(0, 1) < prob + ### aviod mask all, which will cause find_unused_parameters error + if mask.all(): + mask[0]=False + return mask + +class RelativePositionBias(nn.Module): + def __init__( + self, + heads = 8, + num_buckets = 32, + max_distance = 128 + ): + super().__init__() + self.num_buckets = num_buckets + self.max_distance = max_distance + self.relative_attention_bias = nn.Embedding(num_buckets, heads) + + @staticmethod + def _relative_position_bucket(relative_position, num_buckets = 32, max_distance = 128): + ret = 0 + n = -relative_position + + num_buckets //= 2 + ret += (n < 0).long() * num_buckets + n = torch.abs(n) + + max_exact = num_buckets // 2 + is_small = n < max_exact + + val_if_large = max_exact + ( + torch.log(n.float() / max_exact) / math.log(max_distance / max_exact) * (num_buckets - max_exact) + ).long() + val_if_large = torch.min(val_if_large, torch.full_like(val_if_large, num_buckets - 1)) + + ret += torch.where(is_small, n, val_if_large) + return ret + + def forward(self, n, device): + q_pos = torch.arange(n, dtype = torch.long, device = device) + k_pos = torch.arange(n, dtype = torch.long, device = device) + rel_pos = rearrange(k_pos, 'j -> 1 j') - rearrange(q_pos, 'i -> i 1') + rp_bucket = self._relative_position_bucket(rel_pos, num_buckets = self.num_buckets, max_distance = self.max_distance) + values = self.relative_attention_bias(rp_bucket) + return rearrange(values, 'i j h -> h i j') + +class SpatialTransformer(nn.Module): + """ + Transformer block for image-like data. + First, project the input (aka embedding) + and reshape to b, t, d. + Then apply standard transformer action. + Finally, reshape to image + NEW: use_linear for more efficiency instead of the 1x1 convs + """ + def __init__(self, in_channels, n_heads, d_head, + depth=1, dropout=0., context_dim=None, + disable_self_attn=False, use_linear=False, + use_checkpoint=True): + super().__init__() + if exists(context_dim) and not isinstance(context_dim, list): + context_dim = [context_dim] + self.in_channels = in_channels + inner_dim = n_heads * d_head + self.norm = torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True) + if not use_linear: + self.proj_in = nn.Conv2d(in_channels, + inner_dim, + kernel_size=1, + stride=1, + padding=0) + else: + self.proj_in = nn.Linear(in_channels, inner_dim) + + self.transformer_blocks = nn.ModuleList( + [BasicTransformerBlock(inner_dim, n_heads, d_head, dropout=dropout, context_dim=context_dim[d], + disable_self_attn=disable_self_attn, checkpoint=use_checkpoint) + for d in range(depth)] + ) + if not use_linear: + self.proj_out = zero_module(nn.Conv2d(inner_dim, + in_channels, + kernel_size=1, + stride=1, + padding=0)) + else: + self.proj_out = zero_module(nn.Linear(in_channels, inner_dim)) + self.use_linear = use_linear + + def forward(self, x, context=None): + # note: if no context is given, cross-attention defaults to self-attention + if not isinstance(context, list): + context = [context] + b, c, h, w = x.shape + x_in = x + x = self.norm(x) + if not self.use_linear: + x = self.proj_in(x) + x = rearrange(x, 'b c h w -> b (h w) c').contiguous() + if self.use_linear: + x = self.proj_in(x) + for i, block in enumerate(self.transformer_blocks): + x = block(x, context=context[i]) + if self.use_linear: + x = self.proj_out(x) + x = rearrange(x, 'b (h w) c -> b c h w', h=h, w=w).contiguous() + if not self.use_linear: + x = self.proj_out(x) + return x + x_in + +import os +_ATTN_PRECISION = os.environ.get("ATTN_PRECISION", "fp32") + +class CrossAttention(nn.Module): + def __init__(self, query_dim, context_dim=None, heads=8, dim_head=64, dropout=0.): + super().__init__() + inner_dim = dim_head * heads + context_dim = default(context_dim, query_dim) + + self.scale = dim_head ** -0.5 + self.heads = heads + + self.to_q = nn.Linear(query_dim, inner_dim, bias=False) + self.to_k = nn.Linear(context_dim, inner_dim, bias=False) + self.to_v = nn.Linear(context_dim, inner_dim, bias=False) + + self.to_out = nn.Sequential( + nn.Linear(inner_dim, query_dim), + nn.Dropout(dropout) + ) + + def forward(self, x, context=None, mask=None): + h = self.heads + + q = self.to_q(x) + context = default(context, x) + k = self.to_k(context) + v = self.to_v(context) + + q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), (q, k, v)) + + # force cast to fp32 to avoid overflowing + if _ATTN_PRECISION =="fp32": + with torch.autocast(enabled=False, device_type = 'cuda'): + q, k = q.float(), k.float() + sim = torch.einsum('b i d, b j d -> b i j', q, k) * self.scale + else: + sim = torch.einsum('b i d, b j d -> b i j', q, k) * self.scale + + del q, k + + if exists(mask): + mask = rearrange(mask, 'b ... -> b (...)') + max_neg_value = -torch.finfo(sim.dtype).max + mask = repeat(mask, 'b j -> (b h) () j', h=h) + sim.masked_fill_(~mask, max_neg_value) + + # attention, what we cannot get enough of + sim = sim.softmax(dim=-1) + + out = torch.einsum('b i j, b j d -> b i d', sim, v) + out = rearrange(out, '(b h) n d -> b n (h d)', h=h) + return self.to_out(out) + +class BasicTransformerBlock(nn.Module): + def __init__(self, dim, n_heads, d_head, dropout=0., context_dim=None, gated_ff=True, checkpoint=True, + disable_self_attn=False): + super().__init__() + # attn_mode = "softmax-xformers" if XFORMERS_IS_AVAILBLE else "softmax" + # assert attn_mode in self.ATTENTION_MODES + attn_cls = CrossAttention + self.disable_self_attn = disable_self_attn + self.attn1 = attn_cls(query_dim=dim, heads=n_heads, dim_head=d_head, dropout=dropout, + context_dim=context_dim if self.disable_self_attn else None) # is a self-attention if not self.disable_self_attn + self.ff = FeedForward(dim, dropout=dropout, glu=gated_ff) + self.attn2 = attn_cls(query_dim=dim, context_dim=context_dim, + heads=n_heads, dim_head=d_head, dropout=dropout) # is self-attn if context is none + self.norm1 = nn.LayerNorm(dim) + self.norm2 = nn.LayerNorm(dim) + self.norm3 = nn.LayerNorm(dim) + self.checkpoint = checkpoint + + def forward_(self, x, context=None): + return checkpoint(self._forward, (x, context), self.parameters(), self.checkpoint) + + def forward(self, x, context=None): + x = self.attn1(self.norm1(x), context=context if self.disable_self_attn else None) + x + x = self.attn2(self.norm2(x), context=context) + x + x = self.ff(self.norm3(x)) + x + return x + +# feedforward +class GEGLU(nn.Module): + def __init__(self, dim_in, dim_out): + super().__init__() + self.proj = nn.Linear(dim_in, dim_out * 2) + + def forward(self, x): + x, gate = self.proj(x).chunk(2, dim=-1) + return x * F.gelu(gate) + +def zero_module(module): + """ + Zero out the parameters of a module and return it. + """ + for p in module.parameters(): + p.detach().zero_() + return module + +class FeedForward(nn.Module): + def __init__(self, dim, dim_out=None, mult=4, glu=False, dropout=0.): + super().__init__() + inner_dim = int(dim * mult) + dim_out = default(dim_out, dim) + project_in = nn.Sequential( + nn.Linear(dim, inner_dim), + nn.GELU() + ) if not glu else GEGLU(dim, inner_dim) + + self.net = nn.Sequential( + project_in, + nn.Dropout(dropout), + nn.Linear(inner_dim, dim_out) + ) + + def forward(self, x): + return self.net(x) + +class Upsample(nn.Module): + """ + An upsampling layer with an optional convolution. + :param channels: channels in the inputs and outputs. + :param use_conv: a bool determining if a convolution is applied. + :param dims: determines if the signal is 1D, 2D, or 3D. If 3D, then + upsampling occurs in the inner-two dimensions. + """ + + def __init__(self, channels, use_conv, dims=2, out_channels=None, padding=1): + super().__init__() + self.channels = channels + self.out_channels = out_channels or channels + self.use_conv = use_conv + self.dims = dims + if use_conv: + self.conv = nn.Conv2d(self.channels, self.out_channels, 3, padding=padding) + + def forward(self, x): + assert x.shape[1] == self.channels + if self.dims == 3: + x = F.interpolate( + x, (x.shape[2], x.shape[3] * 2, x.shape[4] * 2), mode="nearest" + ) + else: + x = F.interpolate(x, scale_factor=2, mode="nearest") + if self.use_conv: + x = self.conv(x) + return x + + +class ResBlock(nn.Module): + """ + A residual block that can optionally change the number of channels. + :param channels: the number of input channels. + :param emb_channels: the number of timestep embedding channels. + :param dropout: the rate of dropout. + :param out_channels: if specified, the number of out channels. + :param use_conv: if True and out_channels is specified, use a spatial + convolution instead of a smaller 1x1 convolution to change the + channels in the skip connection. + :param dims: determines if the signal is 1D, 2D, or 3D. + :param use_checkpoint: if True, use gradient checkpointing on this module. + :param up: if True, use this block for upsampling. + :param down: if True, use this block for downsampling. + """ + def __init__( + self, + channels, + emb_channels, + dropout, + out_channels=None, + use_conv=False, + use_scale_shift_norm=False, + dims=2, + up=False, + down=False, + use_temporal_conv=True, + use_image_dataset=False, + ): + super().__init__() + self.channels = channels + self.emb_channels = emb_channels + self.dropout = dropout + self.out_channels = out_channels or channels + self.use_conv = use_conv + self.use_scale_shift_norm = use_scale_shift_norm + self.use_temporal_conv = use_temporal_conv + + self.in_layers = nn.Sequential( + nn.GroupNorm(32, channels), + nn.SiLU(), + nn.Conv2d(channels, self.out_channels, 3, padding=1), + ) + + self.updown = up or down + + if up: + self.h_upd = Upsample(channels, False, dims) + self.x_upd = Upsample(channels, False, dims) + elif down: + self.h_upd = Downsample(channels, False, dims) + self.x_upd = Downsample(channels, False, dims) + else: + self.h_upd = self.x_upd = nn.Identity() + + self.emb_layers = nn.Sequential( + nn.SiLU(), + nn.Linear( + emb_channels, + 2 * self.out_channels if use_scale_shift_norm else self.out_channels, + ), + ) + self.out_layers = nn.Sequential( + nn.GroupNorm(32, self.out_channels), + nn.SiLU(), + nn.Dropout(p=dropout), + zero_module( + nn.Conv2d(self.out_channels, self.out_channels, 3, padding=1) + ), + ) + + if self.out_channels == channels: + self.skip_connection = nn.Identity() + elif use_conv: + self.skip_connection = conv_nd( + dims, channels, self.out_channels, 3, padding=1 + ) + else: + self.skip_connection = nn.Conv2d(channels, self.out_channels, 1) + + if self.use_temporal_conv: + self.temopral_conv = TemporalConvBlock_v2(self.out_channels, self.out_channels, dropout=0.1, use_image_dataset=use_image_dataset) + # self.temopral_conv_2 = TemporalConvBlock(self.out_channels, self.out_channels, dropout=0.1, use_image_dataset=use_image_dataset) + + def forward(self, x, emb, batch_size): + """ + Apply the block to a Tensor, conditioned on a timestep embedding. + :param x: an [N x C x ...] Tensor of features. + :param emb: an [N x emb_channels] Tensor of timestep embeddings. + :return: an [N x C x ...] Tensor of outputs. + """ + return self._forward(x, emb, batch_size) + + def _forward(self, x, emb, batch_size): + if self.updown: + in_rest, in_conv = self.in_layers[:-1], self.in_layers[-1] + h = in_rest(x) + h = self.h_upd(h) + x = self.x_upd(x) + h = in_conv(h) + else: + h = self.in_layers(x) + emb_out = self.emb_layers(emb).type(h.dtype) + while len(emb_out.shape) < len(h.shape): + emb_out = emb_out[..., None] + if self.use_scale_shift_norm: + out_norm, out_rest = self.out_layers[0], self.out_layers[1:] + scale, shift = th.chunk(emb_out, 2, dim=1) + h = out_norm(h) * (1 + scale) + shift + h = out_rest(h) + else: + h = h + emb_out + h = self.out_layers(h) + h = self.skip_connection(x) + h + + if self.use_temporal_conv: + h = rearrange(h, '(b f) c h w -> b c f h w', b=batch_size) + h = self.temopral_conv(h) + # h = self.temopral_conv_2(h) + h = rearrange(h, 'b c f h w -> (b f) c h w') + return h + +class Downsample(nn.Module): + """ + A downsampling layer with an optional convolution. + :param channels: channels in the inputs and outputs. + :param use_conv: a bool determining if a convolution is applied. + :param dims: determines if the signal is 1D, 2D, or 3D. If 3D, then + downsampling occurs in the inner-two dimensions. + """ + + def __init__(self, channels, use_conv, dims=2, out_channels=None,padding=1): + super().__init__() + self.channels = channels + self.out_channels = out_channels or channels + self.use_conv = use_conv + self.dims = dims + stride = 2 if dims != 3 else (1, 2, 2) + if use_conv: + self.op = nn.Conv2d(self.channels, self.out_channels, 3, stride=stride, padding=padding) + else: + assert self.channels == self.out_channels + self.op = avg_pool_nd(dims, kernel_size=stride, stride=stride) + + def forward(self, x): + assert x.shape[1] == self.channels + return self.op(x) + +class Resample(nn.Module): + + def __init__(self, in_dim, out_dim, mode): + assert mode in ['none', 'upsample', 'downsample'] + super(Resample, self).__init__() + self.in_dim = in_dim + self.out_dim = out_dim + self.mode = mode + + def forward(self, x, reference=None): + if self.mode == 'upsample': + assert reference is not None + x = F.interpolate(x, size=reference.shape[-2:], mode='nearest') + elif self.mode == 'downsample': + x = F.adaptive_avg_pool2d(x, output_size=tuple(u // 2 for u in x.shape[-2:])) + return x + +class ResidualBlock(nn.Module): + + def __init__(self, in_dim, embed_dim, out_dim, use_scale_shift_norm=True, + mode='none', dropout=0.0): + super(ResidualBlock, self).__init__() + self.in_dim = in_dim + self.embed_dim = embed_dim + self.out_dim = out_dim + self.use_scale_shift_norm = use_scale_shift_norm + self.mode = mode + + # layers + self.layer1 = nn.Sequential( + nn.GroupNorm(32, in_dim), + nn.SiLU(), + nn.Conv2d(in_dim, out_dim, 3, padding=1)) + self.resample = Resample(in_dim, in_dim, mode) + self.embedding = nn.Sequential( + nn.SiLU(), + nn.Linear(embed_dim, out_dim * 2 if use_scale_shift_norm else out_dim)) + self.layer2 = nn.Sequential( + nn.GroupNorm(32, out_dim), + nn.SiLU(), + nn.Dropout(dropout), + nn.Conv2d(out_dim, out_dim, 3, padding=1)) + self.shortcut = nn.Identity() if in_dim == out_dim else nn.Conv2d(in_dim, out_dim, 1) + + # zero out the last layer params + nn.init.zeros_(self.layer2[-1].weight) + + def forward(self, x, e, reference=None): + identity = self.resample(x, reference) + x = self.layer1[-1](self.resample(self.layer1[:-1](x), reference)) + e = self.embedding(e).unsqueeze(-1).unsqueeze(-1).type(x.dtype) + if self.use_scale_shift_norm: + scale, shift = e.chunk(2, dim=1) + x = self.layer2[0](x) * (1 + scale) + shift + x = self.layer2[1:](x) + else: + x = x + e + x = self.layer2(x) + x = x + self.shortcut(identity) + return x + +class AttentionBlock(nn.Module): + + def __init__(self, dim, context_dim=None, num_heads=None, head_dim=None): + # consider head_dim first, then num_heads + num_heads = dim // head_dim if head_dim else num_heads + head_dim = dim // num_heads + assert num_heads * head_dim == dim + super(AttentionBlock, self).__init__() + self.dim = dim + self.context_dim = context_dim + self.num_heads = num_heads + self.head_dim = head_dim + self.scale = math.pow(head_dim, -0.25) + + # layers + self.norm = nn.GroupNorm(32, dim) + self.to_qkv = nn.Conv2d(dim, dim * 3, 1) + if context_dim is not None: + self.context_kv = nn.Linear(context_dim, dim * 2) + self.proj = nn.Conv2d(dim, dim, 1) + + # zero out the last layer params + nn.init.zeros_(self.proj.weight) + + def forward(self, x, context=None): + r"""x: [B, C, H, W]. + context: [B, L, C] or None. + """ + identity = x + b, c, h, w, n, d = *x.size(), self.num_heads, self.head_dim + + # compute query, key, value + x = self.norm(x) + q, k, v = self.to_qkv(x).view(b, n * 3, d, h * w).chunk(3, dim=1) + if context is not None: + ck, cv = self.context_kv(context).reshape(b, -1, n * 2, d).permute(0, 2, 3, 1).chunk(2, dim=1) + k = torch.cat([ck, k], dim=-1) + v = torch.cat([cv, v], dim=-1) + + # compute attention + attn = torch.matmul(q.transpose(-1, -2) * self.scale, k * self.scale) + attn = F.softmax(attn, dim=-1) + + # gather context + x = torch.matmul(v, attn.transpose(-1, -2)) + x = x.reshape(b, c, h, w) + + # output + x = self.proj(x) + return x + identity + + +class TemporalAttentionBlock(nn.Module): + def __init__( + self, + dim, + heads = 4, + dim_head = 32, + rotary_emb = None, + use_image_dataset = False, + use_sim_mask = False + ): + super().__init__() + # consider num_heads first, as pos_bias needs fixed num_heads + # heads = dim // dim_head if dim_head else heads + dim_head = dim // heads + assert heads * dim_head == dim + self.use_image_dataset = use_image_dataset + self.use_sim_mask = use_sim_mask + + self.scale = dim_head ** -0.5 + self.heads = heads + hidden_dim = dim_head * heads + + self.norm = nn.GroupNorm(32, dim) + self.rotary_emb = rotary_emb + self.to_qkv = nn.Linear(dim, hidden_dim * 3)#, bias = False) + self.to_out = nn.Linear(hidden_dim, dim)#, bias = False) + + # nn.init.zeros_(self.to_out.weight) + # nn.init.zeros_(self.to_out.bias) + + def forward( + self, + x, + pos_bias = None, + focus_present_mask = None, + video_mask = None + ): + + identity = x + n, height, device = x.shape[2], x.shape[-2], x.device + + x = self.norm(x) + x = rearrange(x, 'b c f h w -> b (h w) f c') + + qkv = self.to_qkv(x).chunk(3, dim = -1) + + if exists(focus_present_mask) and focus_present_mask.all(): + # if all batch samples are focusing on present + # it would be equivalent to passing that token's values (v=qkv[-1]) through to the output + values = qkv[-1] + out = self.to_out(values) + out = rearrange(out, 'b (h w) f c -> b c f h w', h = height) + + return out + identity + + # split out heads + # q, k, v = rearrange_many(qkv, '... n (h d) -> ... h n d', h = self.heads) + # shape [b (hw) h n c/h], n=f + q= rearrange(qkv[0], '... n (h d) -> ... h n d', h = self.heads) + k= rearrange(qkv[1], '... n (h d) -> ... h n d', h = self.heads) + v= rearrange(qkv[2], '... n (h d) -> ... h n d', h = self.heads) + + + # scale + + q = q * self.scale + + # rotate positions into queries and keys for time attention + if exists(self.rotary_emb): + q = self.rotary_emb.rotate_queries_or_keys(q) + k = self.rotary_emb.rotate_queries_or_keys(k) + + # similarity + # shape [b (hw) h n n], n=f + sim = torch.einsum('... h i d, ... h j d -> ... h i j', q, k) + + # relative positional bias + + if exists(pos_bias): + # print(sim.shape,pos_bias.shape) + sim = sim + pos_bias + + if (focus_present_mask is None and video_mask is not None): + #video_mask: [B, n] + mask = video_mask[:, None, :] * video_mask[:, :, None] # [b,n,n] + mask = mask.unsqueeze(1).unsqueeze(1) #[b,1,1,n,n] + sim = sim.masked_fill(~mask, -torch.finfo(sim.dtype).max) + elif exists(focus_present_mask) and not (~focus_present_mask).all(): + attend_all_mask = torch.ones((n, n), device = device, dtype = torch.bool) + attend_self_mask = torch.eye(n, device = device, dtype = torch.bool) + + mask = torch.where( + rearrange(focus_present_mask, 'b -> b 1 1 1 1'), + rearrange(attend_self_mask, 'i j -> 1 1 1 i j'), + rearrange(attend_all_mask, 'i j -> 1 1 1 i j'), + ) + + sim = sim.masked_fill(~mask, -torch.finfo(sim.dtype).max) + + if self.use_sim_mask: + sim_mask = torch.tril(torch.ones((n, n), device = device, dtype = torch.bool), diagonal=0) + sim = sim.masked_fill(~sim_mask, -torch.finfo(sim.dtype).max) + + # numerical stability + sim = sim - sim.amax(dim = -1, keepdim = True).detach() + attn = sim.softmax(dim = -1) + + # aggregate values + + out = torch.einsum('... h i j, ... h j d -> ... h i d', attn, v) + out = rearrange(out, '... h n d -> ... n (h d)') + out = self.to_out(out) + + out = rearrange(out, 'b (h w) f c -> b c f h w', h = height) + + if self.use_image_dataset: + out = identity + 0*out + else: + out = identity + out + return out + +class TemporalTransformer(nn.Module): + """ + Transformer block for image-like data. + First, project the input (aka embedding) + and reshape to b, t, d. + Then apply standard transformer action. + Finally, reshape to image + """ + def __init__(self, in_channels, n_heads, d_head, + depth=1, dropout=0., context_dim=None, + disable_self_attn=False, use_linear=False, + use_checkpoint=True, only_self_att=True, multiply_zero=False): + super().__init__() + self.multiply_zero = multiply_zero + self.only_self_att = only_self_att + self.use_adaptor = False + if self.only_self_att: + context_dim = None + if not isinstance(context_dim, list): + context_dim = [context_dim] + self.in_channels = in_channels + inner_dim = n_heads * d_head + self.norm = torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True) + if not use_linear: + self.proj_in = nn.Conv1d(in_channels, + inner_dim, + kernel_size=1, + stride=1, + padding=0) + else: + self.proj_in = nn.Linear(in_channels, inner_dim) + if self.use_adaptor: + self.adaptor_in = nn.Linear(frames, frames) + + self.transformer_blocks = nn.ModuleList( + [BasicTransformerBlock(inner_dim, n_heads, d_head, dropout=dropout, context_dim=context_dim[d], + checkpoint=use_checkpoint) + for d in range(depth)] + ) + if not use_linear: + self.proj_out = zero_module(nn.Conv1d(inner_dim, + in_channels, + kernel_size=1, + stride=1, + padding=0)) + else: + self.proj_out = zero_module(nn.Linear(in_channels, inner_dim)) + if self.use_adaptor: + self.adaptor_out = nn.Linear(frames, frames) + self.use_linear = use_linear + + def forward(self, x, context=None): + # note: if no context is given, cross-attention defaults to self-attention + if self.only_self_att: + context = None + if not isinstance(context, list): + context = [context] + b, c, f, h, w = x.shape + x_in = x + x = self.norm(x) + + if not self.use_linear: + x = rearrange(x, 'b c f h w -> (b h w) c f').contiguous() + x = self.proj_in(x) + # [16384, 16, 320] + if self.use_linear: + x = rearrange(x, '(b f) c h w -> b (h w) f c', f=self.frames).contiguous() + x = self.proj_in(x) + + if self.only_self_att: + x = rearrange(x, 'bhw c f -> bhw f c').contiguous() + for i, block in enumerate(self.transformer_blocks): + x = block(x) + x = rearrange(x, '(b hw) f c -> b hw f c', b=b).contiguous() + else: + x = rearrange(x, '(b hw) c f -> b hw f c', b=b).contiguous() + for i, block in enumerate(self.transformer_blocks): + # context[i] = repeat(context[i], '(b f) l con -> b (f r) l con', r=(h*w)//self.frames, f=self.frames).contiguous() + context[i] = rearrange(context[i], '(b f) l con -> b f l con', f=self.frames).contiguous() + # calculate each batch one by one (since number in shape could not greater then 65,535 for some package) + for j in range(b): + context_i_j = repeat(context[i][j], 'f l con -> (f r) l con', r=(h*w)//self.frames, f=self.frames).contiguous() + x[j] = block(x[j], context=context_i_j) + + if self.use_linear: + x = self.proj_out(x) + x = rearrange(x, 'b (h w) f c -> b f c h w', h=h, w=w).contiguous() + if not self.use_linear: + # x = rearrange(x, 'bhw f c -> bhw c f').contiguous() + x = rearrange(x, 'b hw f c -> (b hw) c f').contiguous() + x = self.proj_out(x) + x = rearrange(x, '(b h w) c f -> b c f h w', b=b, h=h, w=w).contiguous() + + if self.multiply_zero: + x = 0.0 * x + x_in + else: + x = x + x_in + return x + +class TemporalAttentionMultiBlock(nn.Module): + def __init__( + self, + dim, + heads=4, + dim_head=32, + rotary_emb=None, + use_image_dataset=False, + use_sim_mask=False, + temporal_attn_times=1, + ): + super().__init__() + self.att_layers = nn.ModuleList( + [TemporalAttentionBlock(dim, heads, dim_head, rotary_emb, use_image_dataset, use_sim_mask) + for _ in range(temporal_attn_times)] + ) + + def forward( + self, + x, + pos_bias = None, + focus_present_mask = None, + video_mask = None + ): + for layer in self.att_layers: + x = layer(x, pos_bias, focus_present_mask, video_mask) + return x + + +class InitTemporalConvBlock(nn.Module): + + def __init__(self, in_dim, out_dim=None, dropout=0.0,use_image_dataset=False): + super(InitTemporalConvBlock, self).__init__() + if out_dim is None: + out_dim = in_dim#int(1.5*in_dim) + self.in_dim = in_dim + self.out_dim = out_dim + self.use_image_dataset = use_image_dataset + + # conv layers + self.conv = nn.Sequential( + nn.GroupNorm(32, out_dim), + nn.SiLU(), + nn.Dropout(dropout), + nn.Conv3d(out_dim, in_dim, (3, 1, 1), padding = (1, 0, 0))) + + # zero out the last layer params,so the conv block is identity + # nn.init.zeros_(self.conv1[-1].weight) + # nn.init.zeros_(self.conv1[-1].bias) + nn.init.zeros_(self.conv[-1].weight) + nn.init.zeros_(self.conv[-1].bias) + + def forward(self, x): + identity = x + x = self.conv(x) + if self.use_image_dataset: + x = identity + 0*x + else: + x = identity + x + return x + +class TemporalConvBlock(nn.Module): + + def __init__(self, in_dim, out_dim=None, dropout=0.0, use_image_dataset= False): + super(TemporalConvBlock, self).__init__() + if out_dim is None: + out_dim = in_dim#int(1.5*in_dim) + self.in_dim = in_dim + self.out_dim = out_dim + self.use_image_dataset = use_image_dataset + + # conv layers + self.conv1 = nn.Sequential( + nn.GroupNorm(32, in_dim), + nn.SiLU(), + nn.Conv3d(in_dim, out_dim, (3, 1, 1), padding = (1, 0, 0))) + self.conv2 = nn.Sequential( + nn.GroupNorm(32, out_dim), + nn.SiLU(), + nn.Dropout(dropout), + nn.Conv3d(out_dim, in_dim, (3, 1, 1), padding = (1, 0, 0))) + + # zero out the last layer params,so the conv block is identity + # nn.init.zeros_(self.conv1[-1].weight) + # nn.init.zeros_(self.conv1[-1].bias) + nn.init.zeros_(self.conv2[-1].weight) + nn.init.zeros_(self.conv2[-1].bias) + + def forward(self, x): + identity = x + x = self.conv1(x) + x = self.conv2(x) + if self.use_image_dataset: + x = identity + 0*x + else: + x = identity + x + return x + +class TemporalConvBlock_v2(nn.Module): + def __init__(self, in_dim, out_dim=None, dropout=0.0, use_image_dataset=False): + super(TemporalConvBlock_v2, self).__init__() + if out_dim is None: + out_dim = in_dim # int(1.5*in_dim) + self.in_dim = in_dim + self.out_dim = out_dim + self.use_image_dataset = use_image_dataset + + # conv layers + self.conv1 = nn.Sequential( + nn.GroupNorm(32, in_dim), + nn.SiLU(), + nn.Conv3d(in_dim, out_dim, (3, 1, 1), padding = (1, 0, 0))) + self.conv2 = nn.Sequential( + nn.GroupNorm(32, out_dim), + nn.SiLU(), + nn.Dropout(dropout), + nn.Conv3d(out_dim, in_dim, (3, 1, 1), padding = (1, 0, 0))) + self.conv3 = nn.Sequential( + nn.GroupNorm(32, out_dim), + nn.SiLU(), + nn.Dropout(dropout), + nn.Conv3d(out_dim, in_dim, (3, 1, 1), padding = (1, 0, 0))) + self.conv4 = nn.Sequential( + nn.GroupNorm(32, out_dim), + nn.SiLU(), + nn.Dropout(dropout), + nn.Conv3d(out_dim, in_dim, (3, 1, 1), padding = (1, 0, 0))) + + # zero out the last layer params,so the conv block is identity + nn.init.zeros_(self.conv4[-1].weight) + nn.init.zeros_(self.conv4[-1].bias) + + def forward(self, x): + identity = x + x = self.conv1(x) + x = self.conv2(x) + x = self.conv3(x) + x = self.conv4(x) + + if self.use_image_dataset: + x = identity + 0.0 * x + else: + x = identity + x + return x + +@UNET.register_class() +class UNetSDUNCLIPvsFPS(nn.Module): + def __init__(self, + in_dim=7, + dim=512, + y_dim=512, + num_tokens=4, + context_dim=512, + out_dim=6, + dim_mult=[1, 2, 3, 4], + num_heads=None, + head_dim=64, + num_res_blocks=3, + attn_scales=[1 / 2, 1 / 4, 1 / 8], + use_scale_shift_norm=True, + dropout=0.1, + default_fps=8, + temporal_attn_times=1, + temporal_attention = True, + use_checkpoint=False, + use_image_dataset=False, + use_sim_mask = False, + training=True, + inpainting=True, + **kwargs): + embed_dim = dim * 4 + num_heads=num_heads if num_heads else dim//32 + super(UNetSDUNCLIPvsFPS, self).__init__() + self.in_dim = in_dim # 4 + self.num_tokens = num_tokens + self.dim = dim # 320 + self.y_dim = y_dim # 1024 + self.context_dim = context_dim # 1024 + self.embed_dim = embed_dim # 1280 + self.out_dim = out_dim # 4 + self.dim_mult = dim_mult # [1, 2, 4, 4] + ### for temporal attention + self.num_heads = num_heads # 8 + ### for spatial attention + self.default_fps = default_fps + self.head_dim = head_dim # 64 + self.num_res_blocks = num_res_blocks # 2 + self.attn_scales = attn_scales # [1.0, 0.5, 0.25] + self.use_scale_shift_norm = use_scale_shift_norm # True + self.temporal_attn_times = temporal_attn_times # 1 + self.temporal_attention = temporal_attention # True + self.use_checkpoint = use_checkpoint # True + self.use_image_dataset = use_image_dataset # False + self.use_sim_mask = use_sim_mask # False + self.training=training # True + self.inpainting = inpainting # True + + use_linear_in_temporal = False + transformer_depth = 1 + disabled_sa = False + # params + enc_dims = [dim * u for u in [1] + dim_mult] # [320, 320, 640, 1280, 1280] + dec_dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]] # [1280, 1280, 1280, 640, 320] + shortcut_dims = [] + scale = 1.0 + + # embeddings + self.time_embed = nn.Sequential( + nn.Linear(dim, embed_dim), # [320,1280] + nn.SiLU(), + nn.Linear(embed_dim, embed_dim)) + + self.context_embedding = nn.Sequential( + nn.Linear(y_dim, embed_dim), + nn.SiLU(), + nn.Linear(embed_dim, context_dim * self.num_tokens)) + + self.fps_embedding = nn.Sequential( + nn.Linear(dim, embed_dim), + nn.SiLU(), + nn.Linear(embed_dim, embed_dim)) + nn.init.zeros_(self.fps_embedding[-1].weight) + nn.init.zeros_(self.fps_embedding[-1].bias) + + if temporal_attention and not USE_TEMPORAL_TRANSFORMER: + self.rotary_emb = RotaryEmbedding(min(32, head_dim)) + self.time_rel_pos_bias = RelativePositionBias(heads = num_heads, max_distance = 32) # realistically will not be able to generate that many frames of video... yet + + # encoder + self.input_blocks = nn.ModuleList() + init_block = nn.ModuleList([nn.Conv2d(self.in_dim, dim, 3, padding=1)]) + ####need an initial temporal attention? + if temporal_attention: + if USE_TEMPORAL_TRANSFORMER: + init_block.append(TemporalTransformer(dim, num_heads, head_dim, depth=transformer_depth, context_dim=context_dim, + disable_self_attn=disabled_sa, use_linear=use_linear_in_temporal, multiply_zero=use_image_dataset)) + else: + init_block.append(TemporalAttentionMultiBlock(dim, num_heads, head_dim, rotary_emb=self.rotary_emb, temporal_attn_times=temporal_attn_times, use_image_dataset=use_image_dataset)) + # elif temporal_conv: + # init_block.append(InitTemporalConvBlock(dim,dropout=dropout,use_image_dataset=use_image_dataset)) + self.input_blocks.append(init_block) + shortcut_dims.append(dim) + for i, (in_dim, out_dim) in enumerate(zip(enc_dims[:-1], enc_dims[1:])): + for j in range(num_res_blocks): + block = nn.ModuleList([ResBlock(in_dim, embed_dim, dropout, out_channels=out_dim, use_scale_shift_norm=False, use_image_dataset=use_image_dataset,)]) + if scale in attn_scales: + block.append( + SpatialTransformer( + out_dim, out_dim // head_dim, head_dim, depth=1, context_dim=self.context_dim, + disable_self_attn=False, use_linear=True + ) + ) + if self.temporal_attention: + if USE_TEMPORAL_TRANSFORMER: + block.append(TemporalTransformer(out_dim, out_dim // head_dim, head_dim, depth=transformer_depth, context_dim=context_dim, + disable_self_attn=disabled_sa, use_linear=use_linear_in_temporal, multiply_zero=use_image_dataset)) + else: + block.append(TemporalAttentionMultiBlock(out_dim, num_heads, head_dim, rotary_emb = self.rotary_emb, use_image_dataset=use_image_dataset, use_sim_mask=use_sim_mask, temporal_attn_times=temporal_attn_times)) + in_dim = out_dim + self.input_blocks.append(block) + shortcut_dims.append(out_dim) + + # downsample + if i != len(dim_mult) - 1 and j == num_res_blocks - 1: + # block = nn.ModuleList([ResidualBlock(out_dim, embed_dim, out_dim, use_scale_shift_norm, 'downsample')]) + downsample = Downsample( + out_dim, True, dims=2, out_channels=out_dim + ) + shortcut_dims.append(out_dim) + scale /= 2.0 + # block.append(TemporalConvBlock(out_dim,dropout=dropout,use_image_dataset=use_image_dataset)) + self.input_blocks.append(downsample) + + self.middle_block = nn.ModuleList([ + ResBlock(out_dim, embed_dim, dropout, use_scale_shift_norm=False, use_image_dataset=use_image_dataset,), + SpatialTransformer( + out_dim, out_dim // head_dim, head_dim, depth=1, context_dim=self.context_dim, + disable_self_attn=False, use_linear=True + )]) + + if self.temporal_attention: + if USE_TEMPORAL_TRANSFORMER: + self.middle_block.append( + TemporalTransformer( + out_dim, out_dim // head_dim, head_dim, depth=transformer_depth, context_dim=context_dim, + disable_self_attn=disabled_sa, use_linear=use_linear_in_temporal, + multiply_zero=use_image_dataset, + ) + ) + else: + self.middle_block.append(TemporalAttentionMultiBlock(out_dim, num_heads, head_dim, rotary_emb = self.rotary_emb, use_image_dataset=use_image_dataset, use_sim_mask=use_sim_mask, temporal_attn_times=temporal_attn_times)) + + self.middle_block.append(ResBlock(out_dim, embed_dim, dropout, use_scale_shift_norm=False)) + + # decoder + self.output_blocks = nn.ModuleList() + for i, (in_dim, out_dim) in enumerate(zip(dec_dims[:-1], dec_dims[1:])): + for j in range(num_res_blocks + 1): + block = nn.ModuleList([ResBlock(in_dim + shortcut_dims.pop(), embed_dim, dropout, out_dim, use_scale_shift_norm=False, use_image_dataset=use_image_dataset, )]) + if scale in attn_scales: + block.append( + SpatialTransformer( + out_dim, out_dim // head_dim, head_dim, depth=1, context_dim=1024, + disable_self_attn=False, use_linear=True + ) + ) + if self.temporal_attention: + if USE_TEMPORAL_TRANSFORMER: + block.append( + TemporalTransformer( + out_dim, out_dim // head_dim, head_dim, depth=transformer_depth, context_dim=context_dim, + disable_self_attn=disabled_sa, use_linear=use_linear_in_temporal, multiply_zero=use_image_dataset + ) + ) + else: + block.append(TemporalAttentionMultiBlock(out_dim, num_heads, head_dim, rotary_emb =self.rotary_emb, use_image_dataset=use_image_dataset, use_sim_mask=use_sim_mask, temporal_attn_times=temporal_attn_times)) + in_dim = out_dim + + # upsample + if i != len(dim_mult) - 1 and j == num_res_blocks: + upsample = Upsample(out_dim, True, dims=2.0, out_channels=out_dim) + scale *= 2.0 + block.append(upsample) + self.output_blocks.append(block) + + # head + self.out = nn.Sequential( + nn.GroupNorm(32, out_dim), + nn.SiLU(), + nn.Conv2d(out_dim, self.out_dim, 3, padding=1)) + + # zero out the last layer params + nn.init.zeros_(self.out[-1].weight) + + def forward(self, + x, + t, + y, + fps=None, + video_mask=None, + focus_present_mask = None, + prob_focus_present = 0., # probability at which a given batch sample will focus on the present (0. is all off, 1. is completely arrested attention across time) + mask_last_frame_num = 0, # mask last frame num + **kwargs + ): + + batch, c, f, h, w= x.shape + device = x.device + self.batch = batch + if fps is None: + fps = torch.tensor([cfg.default_fps] * batch, dtype=torch.long, device=device) + + #### image and video joint training, if mask_last_frame_num is set, prob_focus_present will be ignored + if mask_last_frame_num > 0: + focus_present_mask = None + video_mask[-mask_last_frame_num:] = False + else: + focus_present_mask = default(focus_present_mask, lambda: prob_mask_like((batch,), prob_focus_present, device = device)) # [False, False] + + if self.temporal_attention and not USE_TEMPORAL_TRANSFORMER: + time_rel_pos_bias = self.time_rel_pos_bias(x.shape[2], device = x.device) + else: + time_rel_pos_bias = None + + # embeddings + embeddings = self.time_embed(sinusoidal_embedding(t, self.dim)) + self.fps_embedding(sinusoidal_embedding(fps, self.dim)) + + context = self.context_embedding(y) + context = context.view(-1, self.num_tokens, self.context_dim) + + # repeat f times for spatial e and context + embeddings = embeddings.repeat_interleave(repeats=f, dim=0) + context=context.repeat_interleave(repeats=f,dim=0) + + ## always in shape (b f) c h w, except for temporal layer + x = rearrange(x, 'b c f h w -> (b f) c h w') + # encoder + xs = [] + for block in self.input_blocks: + x = self._forward_single(block, x, embeddings, context, time_rel_pos_bias, focus_present_mask, video_mask) + xs.append(x) + + # middle + for block in self.middle_block: + x = self._forward_single(block, x, embeddings, context, time_rel_pos_bias,focus_present_mask, video_mask) + + # decoder + for block in self.output_blocks: + x = torch.cat([x, xs.pop()], dim=1) + x = self._forward_single(block, x, embeddings, context, time_rel_pos_bias,focus_present_mask, video_mask, reference=xs[-1] if len(xs) > 0 else None) + + # head + x = self.out(x) # [32, 4, 32, 32] + + # reshape back to (b c f h w) + x = rearrange(x, '(b f) c h w -> b c f h w', b = batch) + return x + + def _forward_single(self, module, x, e, context, time_rel_pos_bias, focus_present_mask, video_mask, reference=None): + if isinstance(module, ResidualBlock): + module = checkpoint_wrapper(module) if self.use_checkpoint else module + x = x.contiguous() + x = module(x, e, reference) + elif isinstance(module, ResBlock): + module = checkpoint_wrapper(module) if self.use_checkpoint else module + x = x.contiguous() + x = module(x, e, self.batch) + elif isinstance(module, SpatialTransformer): + module = checkpoint_wrapper(module) if self.use_checkpoint else module + x = module(x, context) + elif isinstance(module, TemporalTransformer): + module = checkpoint_wrapper(module) if self.use_checkpoint else module + x = rearrange(x, '(b f) c h w -> b c f h w', b = self.batch) + x = module(x, context) + x = rearrange(x, 'b c f h w -> (b f) c h w') + elif isinstance(module, CrossAttention): + module = checkpoint_wrapper(module) if self.use_checkpoint else module + x = module(x, context) + elif isinstance(module, BasicTransformerBlock): + module = checkpoint_wrapper(module) if self.use_checkpoint else module + x = module(x, context) + elif isinstance(module, FeedForward): + x = module(x, context) + elif isinstance(module, Upsample): + x = module(x) + elif isinstance(module, Downsample): + x = module(x) + elif isinstance(module, Resample): + x = module(x, reference) + elif isinstance(module, TemporalAttentionBlock): + module = checkpoint_wrapper(module) if self.use_checkpoint else module + x = rearrange(x, '(b f) c h w -> b c f h w', b = self.batch) + x = module(x, time_rel_pos_bias, focus_present_mask, video_mask) + x = rearrange(x, 'b c f h w -> (b f) c h w') + elif isinstance(module, TemporalAttentionMultiBlock): + module = checkpoint_wrapper(module) if self.use_checkpoint else module + x = rearrange(x, '(b f) c h w -> b c f h w', b = self.batch) + x = module(x, time_rel_pos_bias, focus_present_mask, video_mask) + x = rearrange(x, 'b c f h w -> (b f) c h w') + elif isinstance(module, InitTemporalConvBlock): + module = checkpoint_wrapper(module) if self.use_checkpoint else module + x = rearrange(x, '(b f) c h w -> b c f h w', b = self.batch) + x = module(x) + x = rearrange(x, 'b c f h w -> (b f) c h w') + elif isinstance(module, TemporalConvBlock): + module = checkpoint_wrapper(module) if self.use_checkpoint else module + x = rearrange(x, '(b f) c h w -> b c f h w', b = self.batch) + x = module(x) + x = rearrange(x, 'b c f h w -> (b f) c h w') + elif isinstance(module, nn.ModuleList): + for block in module: + x = self._forward_single(block, x, e, context, time_rel_pos_bias, focus_present_mask, video_mask, reference) + else: + x = module(x) + return x diff --git a/modelscope/models/multi_modal/image_to_video/utils/__init__.py b/modelscope/models/multi_modal/image_to_video/utils/__init__.py new file mode 100755 index 00000000..7adfa377 --- /dev/null +++ b/modelscope/models/multi_modal/image_to_video/utils/__init__.py @@ -0,0 +1 @@ +from .registry_class import * \ No newline at end of file diff --git a/modelscope/models/multi_modal/image_to_video/utils/config.py b/modelscope/models/multi_modal/image_to_video/utils/config.py new file mode 100755 index 00000000..24e10a96 --- /dev/null +++ b/modelscope/models/multi_modal/image_to_video/utils/config.py @@ -0,0 +1,166 @@ +import torch +import logging +import os.path as osp +from datetime import datetime +from easydict import EasyDict +import os + +cfg = EasyDict(__name__='Config: VideoLDM Decoder') + +# ---------------------------work dir-------------------------- +cfg.work_dir = 'workspace/' + + +# ---------------------------Global Variable----------------------------------- +cfg.resolution = [448, 256] +# ----------------------------------------------------------------------------- + + +# ---------------------------Dataset Parameter--------------------------------- +cfg.mean = [0.5, 0.5, 0.5] +cfg.std = [0.5, 0.5, 0.5] +cfg.max_words = 1000 + +# PlaceHolder +cfg.vit_out_dim = 1024 +cfg.vit_resolution = [224, 224] #336 +cfg.depth_clamp = 10.0 +cfg.misc_size = 384 +cfg.depth_std = 20.0 + + +cfg.frame_lens = 32 #[32, 32, 32, 1] +cfg.sample_fps = 8 #[4, ] + +cfg.batch_sizes = 1 +# ----------------------------------------------------------------------------- + + +# ---------------------------Mode Parameters----------------------------------- +# Diffusion +cfg.schedule = 'cosine' +cfg.num_timesteps = 1000 +cfg.mean_type = 'v' #'eps' +cfg.var_type = 'fixed_small' # NOTE: to stabilize training and avoid NaN +cfg.loss_type = 'mse' +cfg.ddim_timesteps = 50 # official: 250 +cfg.ddim_eta = 0.0 +cfg.clamp = 1.0 +cfg.share_noise = False +cfg.use_div_loss = False +cfg.noise_strength = 0.1 + +# classifier-free guidance +cfg.p_zero = 0.9 +cfg.guide_scale = 3.0 + +# clip vision encoder +cfg.vit_mean = [0.48145466, 0.4578275, 0.40821073] +cfg.vit_std = [0.26862954, 0.26130258, 0.27577711] + +# Model +cfg.scale_factor = 0.18215 +cfg.use_fp16 = True +cfg.temporal_attention = True +cfg.decoder_bs = 8 + +cfg.UNet = { + 'type': 'UNetSDUNCLIPvsFPS', + 'in_dim': 4, + 'dim': 320, + 'y_dim': cfg.vit_out_dim, + 'context_dim': 1024, + 'out_dim': 8 if cfg.var_type.startswith('learned') else 4, + 'dim_mult': [1, 2, 4, 4], + 'num_heads': 8, + 'head_dim': 64, + 'num_res_blocks': 2, + 'attn_scales': [1 / 1, 1 / 2, 1 / 4], + 'dropout': 0.1, + 'temporal_attention': cfg.temporal_attention, + 'temporal_attn_times': 1, + 'use_checkpoint': False, + 'use_fps_condition': False, + 'use_sim_mask': False, + 'num_tokens': 4, + 'default_fps': 8, + 'input_dim': 1024 +} + +cfg.guidances = [] + +# auotoencoder from stabel diffusion +cfg.auto_encoder = { + 'type': 'AutoencoderKL', + 'ddconfig': { + 'double_z': True, + 'z_channels': 4, + 'resolution': 256, + 'in_channels': 3, + 'out_ch': 3, + 'ch': 128, + 'ch_mult': [1, 2, 4, 4], + 'num_res_blocks': 2, + 'attn_resolutions': [], + 'dropout': 0.0 + }, + 'embed_dim': 4, + 'pretrained': 'v2-1_512-ema-pruned.ckpt' +} +# clip embedder +cfg.embedder = { + 'type': 'FrozenOpenCLIPVisualEmbedder', + 'layer': 'penultimate', + 'vit_resolution': [224, 224], + 'pretrained': 'open_clip_pytorch_model.bin' +} +# ----------------------------------------------------------------------------- + +# ---------------------------Training Settings--------------------------------- +# training and optimizer +cfg.ema_decay = 0.9999 +cfg.num_steps = 600000 +cfg.lr = 5e-5 +cfg.weight_decay = 0.0 +cfg.betas = (0.9, 0.999) +cfg.eps = 1.0e-8 +cfg.chunk_size = 16 +cfg.alpha = 0.7 +cfg.save_ckp_interval = 1000 +# ----------------------------------------------------------------------------- + + +# ----------------------------Pretrain Settings--------------------------------- +## Default: load 2d pretrain +cfg.fix_weight = False +cfg.load_match = False +cfg.pretrained_checkpoint = 'v2-1_512-ema-pruned.ckpt' +cfg.pretrained_image_keys = 'stable_diffusion_image_key_temporal_attention_x1.json' +cfg.resume_checkpoint = "img2video_ldm_0779000.pth" +# ----------------------------------------------------------------------------- + + +# -----------------------------Visual------------------------------------------- +# Visual videos +cfg.viz_interval = 1000 +cfg.visual_train = { + 'type': 'VisualVideoTextDuringTrain', +} +cfg.visual_inference = { + 'type': 'VisualGeneratedVideos', +} +cfg.inference_list_path = '' + +# logging +cfg.log_interval = 100 + +### Default log_dir +cfg.log_dir = "workspace/output_data" #'workspace/videoldms' +# ----------------------------------------------------------------------------- + + +# ---------------------------Others-------------------------------------------- +# seed +cfg.seed = 8888 +# ----------------------------------------------------------------------------- + diff --git a/modelscope/models/multi_modal/image_to_video/utils/diffusion.py b/modelscope/models/multi_modal/image_to_video/utils/diffusion.py new file mode 100755 index 00000000..e86eae7f --- /dev/null +++ b/modelscope/models/multi_modal/image_to_video/utils/diffusion.py @@ -0,0 +1,377 @@ +import torch +import math + + +__all__ = ['GaussianDiffusion', 'beta_schedule'] + +def _i(tensor, t, x): + r"""Index tensor using t and format the output according to x. + """ + shape = (x.size(0), ) + (1, ) * (x.ndim - 1) + if tensor.device != x.device: + tensor = tensor.to(x.device) + return tensor[t].view(shape).to(x) + +def beta_schedule(schedule, num_timesteps=1000, init_beta=None, last_beta=None): + if schedule == 'linear': + scale = 1000.0 / num_timesteps + init_beta = init_beta or scale * 0.0001 + last_beta = last_beta or scale * 0.02 + return torch.linspace(init_beta, last_beta, num_timesteps, dtype=torch.float64) + elif schedule == 'quadratic': + init_beta = init_beta or 0.0015 + last_beta = last_beta or 0.0195 + return torch.linspace(init_beta ** 0.5, last_beta ** 0.5, num_timesteps, dtype=torch.float64) ** 2 + elif schedule == 'cosine': + betas = [] + for step in range(num_timesteps): + t1 = step / num_timesteps + t2 = (step + 1) / num_timesteps + fn = lambda u: math.cos((u + 0.008) / 1.008 * math.pi / 2) ** 2 + betas.append(min(1.0 - fn(t2) / fn(t1), 0.999)) + return torch.tensor(betas, dtype=torch.float64) + else: + raise ValueError(f'Unsupported schedule: {schedule}') + +class GaussianDiffusion(object): + + def __init__(self, + betas, + mean_type='eps', + var_type='learned_range', + loss_type='mse', + epsilon = 1e-12, + rescale_timesteps=False, + noise_strength=0.0): + # check input + if not isinstance(betas, torch.DoubleTensor): + betas = torch.tensor(betas, dtype=torch.float64) + assert min(betas) > 0 and max(betas) <= 1 + assert mean_type in ['x0', 'x_{t-1}', 'eps', 'v'] + assert var_type in ['learned', 'learned_range', 'fixed_large', 'fixed_small'] + assert loss_type in ['mse', 'rescaled_mse', 'kl', 'rescaled_kl', 'l1', 'rescaled_l1','charbonnier'] + self.betas = betas + self.num_timesteps = len(betas) + self.mean_type = mean_type # eps + self.var_type = var_type # 'fixed_small' + self.loss_type = loss_type # mse + self.epsilon = epsilon # 1e-12 + self.rescale_timesteps = rescale_timesteps # False + self.noise_strength = noise_strength # 0.0 + + # alphas + alphas = 1 - self.betas + self.alphas_cumprod = torch.cumprod(alphas, dim=0) + self.alphas_cumprod_prev = torch.cat([alphas.new_ones([1]), self.alphas_cumprod[:-1]]) + self.alphas_cumprod_next = torch.cat([self.alphas_cumprod[1:], alphas.new_zeros([1])]) + + # q(x_t | x_{t-1}) + self.sqrt_alphas_cumprod = torch.sqrt(self.alphas_cumprod) + self.sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - self.alphas_cumprod) + self.log_one_minus_alphas_cumprod = torch.log(1.0 - self.alphas_cumprod) + self.sqrt_recip_alphas_cumprod = torch.sqrt(1.0 / self.alphas_cumprod) + self.sqrt_recipm1_alphas_cumprod = torch.sqrt(1.0 / self.alphas_cumprod - 1) + + # q(x_{t-1} | x_t, x_0) + self.posterior_variance = betas * (1.0 - self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod) + self.posterior_log_variance_clipped = torch.log(self.posterior_variance.clamp(1e-20)) + self.posterior_mean_coef1 = betas * torch.sqrt(self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod) + self.posterior_mean_coef2 = (1.0 - self.alphas_cumprod_prev) * torch.sqrt(alphas) / (1.0 - self.alphas_cumprod) + + def sample_loss(self, x0, noise=None): + if noise is None: + noise = torch.randn_like(x0) + if self.noise_strength > 0: + b, c, f, _, _= x0.shape + offset_noise = torch.randn(b, c, f, 1, 1, device=x0.device) + noise = noise + self.noise_strength * offset_noise + return noise + + def q_sample(self, x0, t, noise=None): + r"""Sample from q(x_t | x_0). + """ + # noise = torch.randn_like(x0) if noise is None else noise + noise = self.sample_loss(x0, noise) + return _i(self.sqrt_alphas_cumprod, t, x0) * x0 + \ + _i(self.sqrt_one_minus_alphas_cumprod, t, x0) * noise + + def q_mean_variance(self, x0, t): + r"""Distribution of q(x_t | x_0). + """ + mu = _i(self.sqrt_alphas_cumprod, t, x0) * x0 + var = _i(1.0 - self.alphas_cumprod, t, x0) + log_var = _i(self.log_one_minus_alphas_cumprod, t, x0) + return mu, var, log_var + + def q_posterior_mean_variance(self, x0, xt, t): + r"""Distribution of q(x_{t-1} | x_t, x_0). + """ + mu = _i(self.posterior_mean_coef1, t, xt) * x0 + _i(self.posterior_mean_coef2, t, xt) * xt + var = _i(self.posterior_variance, t, xt) + log_var = _i(self.posterior_log_variance_clipped, t, xt) + return mu, var, log_var + + @torch.no_grad() + def p_sample(self, xt, t, model, model_kwargs={}, clamp=None, percentile=None, condition_fn=None, guide_scale=None): + r"""Sample from p(x_{t-1} | x_t). + - condition_fn: for classifier-based guidance (guided-diffusion). + - guide_scale: for classifier-free guidance (glide/dalle-2). + """ + # predict distribution of p(x_{t-1} | x_t) + mu, var, log_var, x0 = self.p_mean_variance(xt, t, model, model_kwargs, clamp, percentile, guide_scale) + + # random sample (with optional conditional function) + noise = torch.randn_like(xt) + mask = t.ne(0).float().view(-1, *((1, ) * (xt.ndim - 1))) # no noise when t == 0 + if condition_fn is not None: + grad = condition_fn(xt, self._scale_timesteps(t), **model_kwargs) + mu = mu.float() + var * grad.float() + xt_1 = mu + mask * torch.exp(0.5 * log_var) * noise + return xt_1, x0 + + @torch.no_grad() + def p_sample_loop(self, noise, model, model_kwargs={}, clamp=None, percentile=None, condition_fn=None, guide_scale=None): + r"""Sample from p(x_{t-1} | x_t) p(x_{t-2} | x_{t-1}) ... p(x_0 | x_1). + """ + # prepare input + b = noise.size(0) + xt = noise + + # diffusion process + for step in torch.arange(self.num_timesteps).flip(0): + t = torch.full((b, ), step, dtype=torch.long, device=xt.device) + xt, _ = self.p_sample(xt, t, model, model_kwargs, clamp, percentile, condition_fn, guide_scale) + return xt + + def p_mean_variance(self, xt, t, model, model_kwargs={}, clamp=None, percentile=None, guide_scale=None): + r"""Distribution of p(x_{t-1} | x_t). + """ + # predict distribution + if guide_scale is None: + out = model(xt, self._scale_timesteps(t), **model_kwargs) + else: + # classifier-free guidance + # (model_kwargs[0]: conditional kwargs; model_kwargs[1]: non-conditional kwargs) + assert isinstance(model_kwargs, list) and len(model_kwargs) == 2 + y_out = model(xt, self._scale_timesteps(t), **model_kwargs[0]) + u_out = model(xt, self._scale_timesteps(t), **model_kwargs[1]) + dim = y_out.size(1) if self.var_type.startswith('fixed') else y_out.size(1) // 2 + out = torch.cat([ + u_out[:, :dim] + guide_scale * (y_out[:, :dim] - u_out[:, :dim]), + y_out[:, dim:]], dim=1) # guide_scale=9.0 + + # compute variance + if self.var_type == 'learned': + out, log_var = out.chunk(2, dim=1) + var = torch.exp(log_var) + elif self.var_type == 'learned_range': + out, fraction = out.chunk(2, dim=1) + min_log_var = _i(self.posterior_log_variance_clipped, t, xt) + max_log_var = _i(torch.log(self.betas), t, xt) + fraction = (fraction + 1) / 2.0 + log_var = fraction * max_log_var + (1 - fraction) * min_log_var + var = torch.exp(log_var) + elif self.var_type == 'fixed_large': + var = _i(torch.cat([self.posterior_variance[1:2], self.betas[1:]]), t, xt) + log_var = torch.log(var) + elif self.var_type == 'fixed_small': + var = _i(self.posterior_variance, t, xt) + log_var = _i(self.posterior_log_variance_clipped, t, xt) + + # compute mean and x0 + if self.mean_type == 'x_{t-1}': + mu = out # x_{t-1} + x0 = _i(1.0 / self.posterior_mean_coef1, t, xt) * mu - \ + _i(self.posterior_mean_coef2 / self.posterior_mean_coef1, t, xt) * xt + elif self.mean_type == 'x0': + x0 = out + mu, _, _ = self.q_posterior_mean_variance(x0, xt, t) + elif self.mean_type == 'eps': + x0 = _i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - \ + _i(self.sqrt_recipm1_alphas_cumprod, t, xt) * out + mu, _, _ = self.q_posterior_mean_variance(x0, xt, t) + elif self.mean_type == 'v': + x0 = _i(self.sqrt_alphas_cumprod, t, xt) * xt - \ + _i(self.sqrt_one_minus_alphas_cumprod, t, xt) * out + mu, _, _ = self.q_posterior_mean_variance(x0, xt, t) + + # restrict the range of x0 + if percentile is not None: + assert percentile > 0 and percentile <= 1 # e.g., 0.995 + s = torch.quantile(x0.flatten(1).abs(), percentile, dim=1).clamp_(1.0).view(-1, 1, 1, 1) + x0 = torch.min(s, torch.max(-s, x0)) / s + elif clamp is not None: + x0 = x0.clamp(-clamp, clamp) + return mu, var, log_var, x0 + + @torch.no_grad() + def ddim_sample(self, xt, t, model, model_kwargs={}, clamp=None, percentile=None, condition_fn=None, guide_scale=None, ddim_timesteps=20, eta=0.0): + r"""Sample from p(x_{t-1} | x_t) using DDIM. + - condition_fn: for classifier-based guidance (guided-diffusion). + - guide_scale: for classifier-free guidance (glide/dalle-2). + """ + stride = self.num_timesteps // ddim_timesteps + + # predict distribution of p(x_{t-1} | x_t) + _, _, _, x0 = self.p_mean_variance(xt, t, model, model_kwargs, clamp, percentile, guide_scale) + if condition_fn is not None: + # x0 -> eps + alpha = _i(self.alphas_cumprod, t, xt) + eps = (_i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - x0) / \ + _i(self.sqrt_recipm1_alphas_cumprod, t, xt) + eps = eps - (1 - alpha).sqrt() * condition_fn(xt, self._scale_timesteps(t), **model_kwargs) + + # eps -> x0 + x0 = _i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - \ + _i(self.sqrt_recipm1_alphas_cumprod, t, xt) * eps + + # derive variables + eps = (_i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - x0) / \ + _i(self.sqrt_recipm1_alphas_cumprod, t, xt) + alphas = _i(self.alphas_cumprod, t, xt) + alphas_prev = _i(self.alphas_cumprod, (t - stride).clamp(0), xt) + sigmas = eta * torch.sqrt((1 - alphas_prev) / (1 - alphas) * (1 - alphas / alphas_prev)) + + # random sample + noise = torch.randn_like(xt) + direction = torch.sqrt(1 - alphas_prev - sigmas ** 2) * eps + mask = t.ne(0).float().view(-1, *((1, ) * (xt.ndim - 1))) + xt_1 = torch.sqrt(alphas_prev) * x0 + direction + mask * sigmas * noise + return xt_1, x0 + + @torch.no_grad() + def ddim_sample_loop(self, noise, model, model_kwargs={}, clamp=None, percentile=None, condition_fn=None, guide_scale=None, ddim_timesteps=20, eta=0.0): + # prepare input + b = noise.size(0) + xt = noise + + # diffusion process (TODO: clamp is inaccurate! Consider replacing the stride by explicit prev/next steps) + steps = (1 + torch.arange(0, self.num_timesteps, self.num_timesteps // ddim_timesteps)).clamp(0, self.num_timesteps - 1).flip(0) + for step in steps: + t = torch.full((b, ), step, dtype=torch.long, device=xt.device) + xt, _ = self.ddim_sample(xt, t, model, model_kwargs, clamp, percentile, condition_fn, guide_scale, ddim_timesteps, eta) + return xt + + @torch.no_grad() + def ddim_reverse_sample(self, xt, t, model, model_kwargs={}, clamp=None, percentile=None, guide_scale=None, ddim_timesteps=20): + r"""Sample from p(x_{t+1} | x_t) using DDIM reverse ODE (deterministic). + """ + stride = self.num_timesteps // ddim_timesteps + + # predict distribution of p(x_{t-1} | x_t) + _, _, _, x0 = self.p_mean_variance(xt, t, model, model_kwargs, clamp, percentile, guide_scale) + + # derive variables + eps = (_i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - x0) / \ + _i(self.sqrt_recipm1_alphas_cumprod, t, xt) + alphas_next = _i( + torch.cat([self.alphas_cumprod, self.alphas_cumprod.new_zeros([1])]), + (t + stride).clamp(0, self.num_timesteps), xt) + + # reverse sample + mu = torch.sqrt(alphas_next) * x0 + torch.sqrt(1 - alphas_next) * eps + return mu, x0 + + @torch.no_grad() + def ddim_reverse_sample_loop(self, x0, model, model_kwargs={}, clamp=None, percentile=None, guide_scale=None, ddim_timesteps=20): + # prepare input + b = x0.size(0) + xt = x0 + + # reconstruction steps + steps = torch.arange(0, self.num_timesteps, self.num_timesteps // ddim_timesteps) + for step in steps: + t = torch.full((b, ), step, dtype=torch.long, device=xt.device) + xt, _ = self.ddim_reverse_sample(xt, t, model, model_kwargs, clamp, percentile, guide_scale, ddim_timesteps) + return xt + + @torch.no_grad() + def plms_sample(self, xt, t, model, model_kwargs={}, clamp=None, percentile=None, condition_fn=None, guide_scale=None, plms_timesteps=20): + r"""Sample from p(x_{t-1} | x_t) using PLMS. + - condition_fn: for classifier-based guidance (guided-diffusion). + - guide_scale: for classifier-free guidance (glide/dalle-2). + """ + stride = self.num_timesteps // plms_timesteps + + # function for compute eps + def compute_eps(xt, t): + # predict distribution of p(x_{t-1} | x_t) + _, _, _, x0 = self.p_mean_variance(xt, t, model, model_kwargs, clamp, percentile, guide_scale) + + # condition + if condition_fn is not None: + # x0 -> eps + alpha = _i(self.alphas_cumprod, t, xt) + eps = (_i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - x0) / \ + _i(self.sqrt_recipm1_alphas_cumprod, t, xt) + eps = eps - (1 - alpha).sqrt() * condition_fn(xt, self._scale_timesteps(t), **model_kwargs) + + # eps -> x0 + x0 = _i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - \ + _i(self.sqrt_recipm1_alphas_cumprod, t, xt) * eps + + # derive eps + eps = (_i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - x0) / \ + _i(self.sqrt_recipm1_alphas_cumprod, t, xt) + return eps + + # function for compute x_0 and x_{t-1} + def compute_x0(eps, t): + # eps -> x0 + x0 = _i(self.sqrt_recip_alphas_cumprod, t, xt) * xt - \ + _i(self.sqrt_recipm1_alphas_cumprod, t, xt) * eps + + # deterministic sample + alphas_prev = _i(self.alphas_cumprod, (t - stride).clamp(0), xt) + direction = torch.sqrt(1 - alphas_prev) * eps + mask = t.ne(0).float().view(-1, *((1, ) * (xt.ndim - 1))) + xt_1 = torch.sqrt(alphas_prev) * x0 + direction + return xt_1, x0 + + # PLMS sample + eps = compute_eps(xt, t) + if len(eps_cache) == 0: + # 2nd order pseudo improved Euler + xt_1, x0 = compute_x0(eps, t) + eps_next = compute_eps(xt_1, (t - stride).clamp(0)) + eps_prime = (eps + eps_next) / 2.0 + elif len(eps_cache) == 1: + # 2nd order pseudo linear multistep (Adams-Bashforth) + eps_prime = (3 * eps - eps_cache[-1]) / 2.0 + elif len(eps_cache) == 2: + # 3nd order pseudo linear multistep (Adams-Bashforth) + eps_prime = (23 * eps - 16 * eps_cache[-1] + 5 * eps_cache[-2]) / 12.0 + elif len(eps_cache) >= 3: + # 4nd order pseudo linear multistep (Adams-Bashforth) + eps_prime = (55 * eps - 59 * eps_cache[-1] + 37 * eps_cache[-2] - 9 * eps_cache[-3]) / 24.0 + xt_1, x0 = compute_x0(eps_prime, t) + return xt_1, x0, eps + + @torch.no_grad() + def plms_sample_loop(self, noise, model, model_kwargs={}, clamp=None, percentile=None, condition_fn=None, guide_scale=None, plms_timesteps=20): + # prepare input + b = noise.size(0) + xt = noise + + # diffusion process + steps = (1 + torch.arange(0, self.num_timesteps, self.num_timesteps // plms_timesteps)).clamp(0, self.num_timesteps - 1).flip(0) + eps_cache = [] + for step in steps: + # PLMS sampling step + t = torch.full((b, ), step, dtype=torch.long, device=xt.device) + xt, _, eps = self.plms_sample(xt, t, model, model_kwargs, clamp, percentile, condition_fn, guide_scale, plms_timesteps, eps_cache) + + # update eps cache + eps_cache.append(eps) + if len(eps_cache) >= 4: + eps_cache.pop(0) + return xt + + + + + def _scale_timesteps(self, t): + if self.rescale_timesteps: + return t.float() * 1000.0 / self.num_timesteps + return t + #return t.float() diff --git a/modelscope/models/multi_modal/image_to_video/utils/registry.py b/modelscope/models/multi_modal/image_to_video/utils/registry.py new file mode 100755 index 00000000..50e399bc --- /dev/null +++ b/modelscope/models/multi_modal/image_to_video/utils/registry.py @@ -0,0 +1,155 @@ +# Copyright 2021 Alibaba Group Holding Limited. All Rights Reserved. + +# Registry class & build_from_config function partially modified from +# https://github.com/open-mmlab/mmcv/blob/master/mmcv/utils/registry.py +# Copyright 2018-2020 Open-MMLab. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import copy +import inspect +import warnings + + +def build_from_config(cfg, registry, **kwargs): + """ Default builder function. + + Args: + cfg (dict): A dict which contains parameters passes to target class or function. + Must contains key 'type', indicates the target class or function name. + registry (Registry): An registry to search target class or function. + kwargs (dict, optional): Other params not in config dict. + + Returns: + Target class object or object returned by invoking function. + + Raises: + TypeError: + KeyError: + Exception: + """ + if not isinstance(cfg, dict): + raise TypeError(f"config must be type dict, got {type(cfg)}") + if "type" not in cfg: + raise KeyError(f"config must contain key type, got {cfg}") + if not isinstance(registry, Registry): + raise TypeError(f"registry must be type Registry, got {type(registry)}") + + cfg = copy.deepcopy(cfg) + + req_type = cfg.pop("type") + req_type_entry = req_type + if isinstance(req_type, str): + req_type_entry = registry.get(req_type) + if req_type_entry is None: + raise KeyError(f"{req_type} not found in {registry.name} registry") + + if kwargs is not None: + cfg.update(kwargs) + + if inspect.isclass(req_type_entry): + try: + return req_type_entry(**cfg) + except Exception as e: + raise Exception(f"Failed to init class {req_type_entry}, with {e}") + elif inspect.isfunction(req_type_entry): + try: + return req_type_entry(**cfg) + except Exception as e: + raise Exception(f"Failed to invoke function {req_type_entry}, with {e}") + else: + raise TypeError(f"type must be str or class, got {type(req_type_entry)}") + + +class Registry(object): + """ A registry maps key to classes or functions. + + Example: + >>> MODELS = Registry('MODELS') + >>> @MODELS.register_class() + >>> class ResNet(object): + >>> pass + >>> resnet = MODELS.build(dict(type="ResNet")) + >>> + >>> import torchvision + >>> @MODELS.register_function("InceptionV3") + >>> def get_inception_v3(pretrained=False, progress=True): + >>> return torchvision.models.inception_v3(pretrained=pretrained, progress=progress) + >>> inception_v3 = MODELS.build(dict(type='InceptionV3', pretrained=True)) + + Args: + name (str): Registry name. + build_func (func, None): Instance construct function. Default is build_from_config. + allow_types (tuple): Indicates how to construct the instance, by constructing class or invoking function. + """ + + def __init__(self, name, build_func=None, allow_types=("class", "function")): + self.name = name + self.allow_types = allow_types + self.class_map = {} + self.func_map = {} + self.build_func = build_func or build_from_config + + def get(self, req_type): + return self.class_map.get(req_type) or self.func_map.get(req_type) + + def build(self, *args, **kwargs): + return self.build_func(*args, **kwargs, registry=self) + + def register_class(self, name=None): + def _register(cls): + if not inspect.isclass(cls): + raise TypeError(f"Module must be type class, got {type(cls)}") + if "class" not in self.allow_types: + raise TypeError(f"Register {self.name} only allows type {self.allow_types}, got class") + module_name = name or cls.__name__ + if module_name in self.class_map: + warnings.warn(f"Class {module_name} already registered by {self.class_map[module_name]}, " + f"will be replaced by {cls}") + self.class_map[module_name] = cls + return cls + + return _register + + def register_function(self, name=None): + def _register(func): + if not inspect.isfunction(func): + raise TypeError(f"Registry must be type function, got {type(func)}") + if "function" not in self.allow_types: + raise TypeError(f"Registry {self.name} only allows type {self.allow_types}, got function") + func_name = name or func.__name__ + if func_name in self.class_map: + warnings.warn(f"Function {func_name} already registered by {self.func_map[func_name]}, " + f"will be replaced by {func}") + self.func_map[func_name] = func + return func + + return _register + + def _list(self): + keys = sorted(list(self.class_map.keys()) + list(self.func_map.keys())) + descriptions = [] + for key in keys: + if key in self.class_map: + descriptions.append(f"{key}: {self.class_map[key]}") + else: + descriptions.append( + f"{key}: ") + return "\n".join(descriptions) + + def __repr__(self): + description = self._list() + description = '\n'.join(['\t' + s for s in description.split('\n')]) + return f"{self.__class__.__name__} [{self.name}], \n" + description + + diff --git a/modelscope/models/multi_modal/image_to_video/utils/registry_class/__init__.py b/modelscope/models/multi_modal/image_to_video/utils/registry_class/__init__.py new file mode 100755 index 00000000..e91472ce --- /dev/null +++ b/modelscope/models/multi_modal/image_to_video/utils/registry_class/__init__.py @@ -0,0 +1,4 @@ +from .model import UNET +from .autoencoder import AUTO_ENCODER +from .embedder import EMBEDDER +from .distrubution import DISTRIBUTION \ No newline at end of file diff --git a/modelscope/models/multi_modal/image_to_video/utils/registry_class/autoencoder.py b/modelscope/models/multi_modal/image_to_video/utils/registry_class/autoencoder.py new file mode 100755 index 00000000..4593b14f --- /dev/null +++ b/modelscope/models/multi_modal/image_to_video/utils/registry_class/autoencoder.py @@ -0,0 +1,11 @@ +# Copyright 2021 Alibaba Group Holding Limited. All Rights Reserved. + +from ..registry import Registry, build_from_config + +def build_autoencoder(cfg, registry, **kwargs): + """ + Except for ordinal autoencoder config, if passing a list of dataset config, then return the concat type of it + """ + return build_from_config(cfg, registry, **kwargs) + +AUTO_ENCODER = Registry("AUTO_ENCODER", build_func=build_autoencoder) \ No newline at end of file diff --git a/modelscope/models/multi_modal/image_to_video/utils/registry_class/distrubution.py b/modelscope/models/multi_modal/image_to_video/utils/registry_class/distrubution.py new file mode 100755 index 00000000..12737150 --- /dev/null +++ b/modelscope/models/multi_modal/image_to_video/utils/registry_class/distrubution.py @@ -0,0 +1,11 @@ +# Copyright 2021 Alibaba Group Holding Limited. All Rights Reserved. + +from ..registry import Registry, build_from_config + +def build_distribution(cfg, registry, **kwargs): + """ + Except for ordinal autoencoder config, if passing a list of dataset config, then return the concat type of it + """ + return build_from_config(cfg, registry, **kwargs) + +DISTRIBUTION = Registry("DISTRIBUTION", build_func=build_distribution) \ No newline at end of file diff --git a/modelscope/models/multi_modal/image_to_video/utils/registry_class/embedder.py b/modelscope/models/multi_modal/image_to_video/utils/registry_class/embedder.py new file mode 100755 index 00000000..48a90adf --- /dev/null +++ b/modelscope/models/multi_modal/image_to_video/utils/registry_class/embedder.py @@ -0,0 +1,11 @@ +# Copyright 2021 Alibaba Group Holding Limited. All Rights Reserved. + +from ..registry import Registry, build_from_config + +def build_embedder(cfg, registry, **kwargs): + """ + Except for ordinal UNet config, if passing a list of dataset config, then return the concat type of it + """ + return build_from_config(cfg, registry, **kwargs) + +EMBEDDER = Registry("EMBEDDER", build_func=build_embedder) diff --git a/modelscope/models/multi_modal/image_to_video/utils/registry_class/model.py b/modelscope/models/multi_modal/image_to_video/utils/registry_class/model.py new file mode 100755 index 00000000..ca9c7b40 --- /dev/null +++ b/modelscope/models/multi_modal/image_to_video/utils/registry_class/model.py @@ -0,0 +1,11 @@ +# Copyright 2021 Alibaba Group Holding Limited. All Rights Reserved. + +from ..registry import Registry, build_from_config + +def build_model(cfg, registry, **kwargs): + """ + Except for ordinal UNet config, if passing a list of dataset config, then return the concat type of it + """ + return build_from_config(cfg, registry, **kwargs) + +UNET = Registry("UNET", build_func=build_model) diff --git a/modelscope/models/multi_modal/image_to_video/utils/seed.py b/modelscope/models/multi_modal/image_to_video/utils/seed.py new file mode 100755 index 00000000..3bd908db --- /dev/null +++ b/modelscope/models/multi_modal/image_to_video/utils/seed.py @@ -0,0 +1,12 @@ +import torch +import random +import numpy as np + + +def setup_seed(seed): + torch.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + np.random.seed(seed) + random.seed(seed) + torch.backends.cudnn.deterministic = True + diff --git a/modelscope/models/multi_modal/image_to_video/utils/shedule.py b/modelscope/models/multi_modal/image_to_video/utils/shedule.py new file mode 100644 index 00000000..8dac53cb --- /dev/null +++ b/modelscope/models/multi_modal/image_to_video/utils/shedule.py @@ -0,0 +1,38 @@ +import torch + + +def beta_schedule(schedule, num_timesteps=1000, init_beta=None, last_beta=None): + ''' + This code defines a function beta_schedule that generates a sequence of beta values based on the given input parameters. These beta values can be used in video diffusion processes. The function has the following parameters: + schedule(str): Determines the type of beta schedule to be generated. It can be 'linear', 'linear_sd', 'quadratic', or 'cosine'. + num_timesteps(int, optional): The number of timesteps for the generated beta schedule. Default is 1000. + init_beta(float, optional): The initial beta value. If not provided, a default value is used based on the chosen schedule. + last_beta(float, optional): The final beta value. If not provided, a default value is used based on the chosen schedule. + The function returns a PyTorch tensor containing the generated beta values. The beta schedule is determined by the schedule parameter: + 1.Linear: Generates a linear sequence of beta values betweeninit_betaandlast_beta. + 2.Linear_sd: Generates a linear sequence of beta values between the square root of init_beta and the square root oflast_beta, and then squares the result. + 3.Quadratic: Similar to the 'linear_sd' schedule, but with different default values forinit_betaandlast_beta. + 4.Cosine: Generates a sequence of beta values based on a cosine function, ensuring the values are between 0 and 0.999. + If an unsupported schedule is provided, a ValueError is raised with a message indicating the issue. + ''' + if schedule == 'linear': + scale = 1000.0 / num_timesteps + init_beta = init_beta or scale * 0.0001 + last_beta = last_beta or scale * 0.02 + return torch.linspace(init_beta, last_beta, num_timesteps, dtype=torch.float64) + elif schedule == 'linear_sd': + return torch.linspace(init_beta ** 0.5, last_beta ** 0.5, num_timesteps, dtype=torch.float64) ** 2 + elif schedule == 'quadratic': + init_beta = init_beta or 0.0015 + last_beta = last_beta or 0.0195 + return torch.linspace(init_beta ** 0.5, last_beta ** 0.5, num_timesteps, dtype=torch.float64) ** 2 + elif schedule == 'cosine': + betas = [] + for step in range(num_timesteps): + t1 = step / num_timesteps + t2 = (step + 1) / num_timesteps + fn = lambda u: math.cos((u + 0.008) / 1.008 * math.pi / 2) ** 2 + betas.append(min(1.0 - fn(t2) / fn(t1), 0.999)) + return torch.tensor(betas, dtype=torch.float64) + else: + raise ValueError(f'Unsupported schedule: {schedule}') \ No newline at end of file diff --git a/modelscope/models/multi_modal/image_to_video/utils/transforms.py b/modelscope/models/multi_modal/image_to_video/utils/transforms.py new file mode 100755 index 00000000..2e2d8ecc --- /dev/null +++ b/modelscope/models/multi_modal/image_to_video/utils/transforms.py @@ -0,0 +1,377 @@ +import torch +import torchvision.transforms.functional as F +import random +import math +import numpy as np +from PIL import Image, ImageFilter + +__all__ = ['Compose', 'Resize', 'Rescale', 'CenterCrop', 'CenterCropV2', 'CenterCropWide', 'RandomCrop', 'RandomCropV2', 'RandomHFlip',\ + 'GaussianBlur', 'ColorJitter', 'RandomGray', 'ToTensor', 'Normalize', "ResizeRandomCrop", "ExtractResizeRandomCrop", "ExtractResizeAssignCrop"] + +# class Compose(object): + +# def __init__(self, transforms): +# self.transforms = transforms + +# def __call__(self, rgb): +# for t in self.transforms: +# rgb = t(rgb) +# return rgb +class Compose(object): + + def __init__(self, transforms): + self.transforms = transforms + + def __getitem__(self, index): + if isinstance(index, slice): + return Compose(self.transforms[index]) + else: + return self.transforms[index] + + def __len__(self): + return len(self.transforms) + + def __call__(self, rgb): + for t in self.transforms: + rgb = t(rgb) + return rgb + +class Resize(object): + + def __init__(self, size=256): + if isinstance(size, int): + size = (size, size) + self.size = size + + def __call__(self, rgb): + if isinstance(rgb, list): + rgb = [u.resize(self.size, Image.BILINEAR) for u in rgb] + else: + rgb = rgb.resize(self.size, Image.BILINEAR) + return rgb + +class Rescale(object): + + def __init__(self, size=256, interpolation=Image.BILINEAR): + self.size = size + self.interpolation = interpolation + + def __call__(self, rgb): + w, h = rgb[0].size + scale = self.size / min(w, h) + out_w, out_h = int(round(w * scale)), int(round(h * scale)) + rgb = [u.resize((out_w, out_h), self.interpolation) for u in rgb] + return rgb + +class CenterCrop(object): + + def __init__(self, size=224): + self.size = size + + def __call__(self, rgb): + w, h = rgb[0].size + assert min(w, h) >= self.size + x1 = (w - self.size) // 2 + y1 = (h - self.size) // 2 + rgb = [u.crop((x1, y1, x1 + self.size, y1 + self.size)) for u in rgb] + return rgb + +class ResizeRandomCrop(object): + + def __init__(self, size=256, size_short=292): + self.size = size + # self.min_area = min_area + self.size_short = size_short + + def __call__(self, rgb): + + # consistent crop between rgb and m + while min(rgb[0].size) >= 2 * self.size_short: + rgb = [u.resize((u.width // 2, u.height // 2), resample=Image.BOX) for u in rgb] + scale = self.size_short / min(rgb[0].size) + rgb = [u.resize((round(scale * u.width), round(scale * u.height)), resample=Image.BICUBIC) for u in rgb] + out_w = self.size + out_h = self.size + w, h = rgb[0].size # (518, 292) + x1 = random.randint(0, w - out_w) + y1 = random.randint(0, h - out_h) + + rgb = [u.crop((x1, y1, x1 + out_w, y1 + out_h)) for u in rgb] + # rgb = [u.resize((self.size, self.size), Image.BILINEAR) for u in rgb] + # # center crop + # x1 = (img[0].width - self.size) // 2 + # y1 = (img[0].height - self.size) // 2 + # img = [u.crop((x1, y1, x1 + self.size, y1 + self.size)) for u in img] + return rgb + + + +class ExtractResizeRandomCrop(object): + + def __init__(self, size=256, size_short=292): + self.size = size + # self.min_area = min_area + self.size_short = size_short + + def __call__(self, rgb): + + # consistent crop between rgb and m + while min(rgb[0].size) >= 2 * self.size_short: + rgb = [u.resize((u.width // 2, u.height // 2), resample=Image.BOX) for u in rgb] + scale = self.size_short / min(rgb[0].size) + rgb = [u.resize((round(scale * u.width), round(scale * u.height)), resample=Image.BICUBIC) for u in rgb] + out_w = self.size + out_h = self.size + w, h = rgb[0].size # (518, 292) + x1 = random.randint(0, w - out_w) + y1 = random.randint(0, h - out_h) + + rgb = [u.crop((x1, y1, x1 + out_w, y1 + out_h)) for u in rgb] + # rgb = [u.resize((self.size, self.size), Image.BILINEAR) for u in rgb] + # # center crop + # x1 = (img[0].width - self.size) // 2 + # y1 = (img[0].height - self.size) // 2 + # img = [u.crop((x1, y1, x1 + self.size, y1 + self.size)) for u in img] + wh = [x1, y1, x1 + out_w, y1 + out_h] + return rgb, wh + + + + +class ExtractResizeAssignCrop(object): + + def __init__(self, size=256, size_short=292): + self.size = size + # self.min_area = min_area + self.size_short = size_short + + def __call__(self, rgb, wh): + + # consistent crop between rgb and m + while min(rgb[0].size) >= 2 * self.size_short: + rgb = [u.resize((u.width // 2, u.height // 2), resample=Image.BOX) for u in rgb] + scale = self.size_short / min(rgb[0].size) + rgb = [u.resize((round(scale * u.width), round(scale * u.height)), resample=Image.BICUBIC) for u in rgb] + # out_w = self.size + # out_h = self.size + # w, h = rgb[0].size # (518, 292) + # x1 = random.randint(0, w - out_w) + # y1 = random.randint(0, h - out_h) + + rgb = [u.crop(wh) for u in rgb] + rgb = [u.resize((self.size, self.size), Image.BILINEAR) for u in rgb] + # # center crop + # x1 = (img[0].width - self.size) // 2 + # y1 = (img[0].height - self.size) // 2 + # img = [u.crop((x1, y1, x1 + self.size, y1 + self.size)) for u in img] + # wh = [x1, y1, x1 + out_w, y1 + out_h] + return rgb + +class CenterCropV2(object): + def __init__(self, size): + self.size = size + + def __call__(self, img): + # fast resize + while min(img[0].size) >= 2 * self.size: + img = [u.resize((u.width // 2, u.height // 2), resample=Image.BOX) for u in img] + scale = self.size / min(img[0].size) + img = [u.resize((round(scale * u.width), round(scale * u.height)), resample=Image.BICUBIC) for u in img] + + # center crop + x1 = (img[0].width - self.size) // 2 + y1 = (img[0].height - self.size) // 2 + img = [u.crop((x1, y1, x1 + self.size, y1 + self.size)) for u in img] + return img + + +class CenterCropWide(object): + def __init__(self, size): + self.size = size + + def __call__(self, img): + if isinstance(img, list): + scale = min(img[0].size[0]/self.size[0], img[0].size[1]/self.size[1]) + img = [u.resize((round(u.width // scale), round(u.height // scale)), resample=Image.BOX) for u in img] + + # center crop + x1 = (img[0].width - self.size[0]) // 2 + y1 = (img[0].height - self.size[1]) // 2 + img = [u.crop((x1, y1, x1 + self.size[0], y1 + self.size[1])) for u in img] + return img + else: + scale = min(img.size[0]/self.size[0], img.size[1]/self.size[1]) + img = img.resize((round(img.width // scale), round(img.height // scale)), resample=Image.BOX) + x1 = (img.width - self.size[0]) // 2 + y1 = (img.height - self.size[1]) // 2 + img = img.crop((x1, y1, x1 + self.size[0], y1 + self.size[1])) + return img + + + +class RandomCrop(object): + + def __init__(self, size=224, min_area=0.4): + self.size = size + self.min_area = min_area + + def __call__(self, rgb): + + # consistent crop between rgb and m + w, h = rgb[0].size + area = w * h + out_w, out_h = float('inf'), float('inf') + while out_w > w or out_h > h: + target_area = random.uniform(self.min_area, 1.0) * area + aspect_ratio = random.uniform(3. / 4., 4. / 3.) + out_w = int(round(math.sqrt(target_area * aspect_ratio))) + out_h = int(round(math.sqrt(target_area / aspect_ratio))) + x1 = random.randint(0, w - out_w) + y1 = random.randint(0, h - out_h) + + rgb = [u.crop((x1, y1, x1 + out_w, y1 + out_h)) for u in rgb] + rgb = [u.resize((self.size, self.size), Image.BILINEAR) for u in rgb] + + return rgb + +class RandomCropV2(object): + + def __init__(self, size=224, min_area=0.4, ratio=(3. / 4., 4. / 3.)): + if isinstance(size, (tuple, list)): + self.size = size + else: + self.size = (size, size) + self.min_area = min_area + self.ratio = ratio + + def _get_params(self, img): + width, height = img.size + area = height * width + + for _ in range(10): + target_area = random.uniform(self.min_area, 1.0) * area + log_ratio = (math.log(self.ratio[0]), math.log(self.ratio[1])) + aspect_ratio = math.exp(random.uniform(*log_ratio)) + + w = int(round(math.sqrt(target_area * aspect_ratio))) + h = int(round(math.sqrt(target_area / aspect_ratio))) + + if 0 < w <= width and 0 < h <= height: + i = random.randint(0, height - h) + j = random.randint(0, width - w) + return i, j, h, w + + # Fallback to central crop + in_ratio = float(width) / float(height) + if (in_ratio < min(self.ratio)): + w = width + h = int(round(w / min(self.ratio))) + elif (in_ratio > max(self.ratio)): + h = height + w = int(round(h * max(self.ratio))) + else: # whole image + w = width + h = height + i = (height - h) // 2 + j = (width - w) // 2 + return i, j, h, w + + def __call__(self, rgb): + i, j, h, w = self._get_params(rgb[0]) + rgb = [F.resized_crop(u, i, j, h, w, self.size) for u in rgb] + return rgb + +class RandomHFlip(object): + + def __init__(self, p=0.5): + self.p = p + + def __call__(self, rgb): + if random.random() < self.p: + rgb = [u.transpose(Image.FLIP_LEFT_RIGHT) for u in rgb] + return rgb + +class GaussianBlur(object): + + def __init__(self, sigmas=[0.1, 2.0], p=0.5): + self.sigmas = sigmas + self.p = p + + def __call__(self, rgb): + if random.random() < self.p: + sigma = random.uniform(*self.sigmas) + rgb = [u.filter(ImageFilter.GaussianBlur(radius=sigma)) for u in rgb] + return rgb + +class ColorJitter(object): + + def __init__(self, brightness=0.4, contrast=0.4, saturation=0.4, hue=0.1, p=0.5): + self.brightness = brightness + self.contrast = contrast + self.saturation = saturation + self.hue = hue + self.p = p + + def __call__(self, rgb): + if random.random() < self.p: + brightness, contrast, saturation, hue = self._random_params() + transforms = [ + lambda f: F.adjust_brightness(f, brightness), + lambda f: F.adjust_contrast(f, contrast), + lambda f: F.adjust_saturation(f, saturation), + lambda f: F.adjust_hue(f, hue)] + random.shuffle(transforms) + for t in transforms: + rgb = [t(u) for u in rgb] + + return rgb + + def _random_params(self): + brightness = random.uniform( + max(0, 1 - self.brightness), 1 + self.brightness) + contrast = random.uniform( + max(0, 1 - self.contrast), 1 + self.contrast) + saturation = random.uniform( + max(0, 1 - self.saturation), 1 + self.saturation) + hue = random.uniform(-self.hue, self.hue) + return brightness, contrast, saturation, hue + +class RandomGray(object): + + def __init__(self, p=0.2): + self.p = p + + def __call__(self, rgb): + if random.random() < self.p: + rgb = [u.convert('L').convert('RGB') for u in rgb] + return rgb + +class ToTensor(object): + + def __call__(self, rgb): + if isinstance(rgb, list): + rgb = torch.stack([F.to_tensor(u) for u in rgb], dim=0) + else: + rgb = F.to_tensor(rgb) + + return rgb + +class Normalize(object): + + def __init__(self, mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]): + self.mean = mean + self.std = std + + def __call__(self, rgb): + rgb = rgb.clone() + rgb.clamp_(0, 1) + if not isinstance(self.mean, torch.Tensor): + self.mean = rgb.new_tensor(self.mean).view(-1) + if not isinstance(self.std, torch.Tensor): + self.std = rgb.new_tensor(self.std).view(-1) + if rgb.dim() == 4: + rgb.sub_(self.mean.view(1, -1, 1, 1)).div_(self.std.view(1, -1, 1, 1)) + elif rgb.dim() == 3: + rgb.sub_(self.mean.view(-1, 1, 1)).div_(self.std.view(-1, 1, 1)) + return rgb + diff --git a/modelscope/pipelines/multi_modal/image_to_video_pipeline.py b/modelscope/pipelines/multi_modal/image_to_video_pipeline.py new file mode 100644 index 00000000..3285854c --- /dev/null +++ b/modelscope/pipelines/multi_modal/image_to_video_pipeline.py @@ -0,0 +1,74 @@ +# Copyright (c) Alibaba, Inc. and its affiliates. + +import os +import tempfile +from typing import Any, Dict, Optional + +import cv2 +import torch +from einops import rearrange + +from modelscope.metainfo import Pipelines +from modelscope.outputs import OutputKeys +from modelscope.pipelines.base import Input, Model, Pipeline +from modelscope.pipelines.builder import PIPELINES +from modelscope.utils.constant import Tasks +from modelscope.utils.logger import get_logger + +logger = get_logger() + +@PIPELINES.register_module(Tasks.image_to_video_task, module_name=Pipelines.image_to_video_task_pipeline) +class ImageToVideoPipeline(Pipeline): + def __init__(self, model: str, **kwargs): + """ + use `model` to create a kws pipeline for prediction + Args: + model: model id on modelscope hub. + """ + super().__init__(model=model, **kwargs) + + def preprocess(self, input: Input, **preprocess_params) -> Dict[str, Any]: + return {'img_path': input['img_path']} + + def forward(self, input: Dict[str, Any], **forward_params) -> Dict[str, Any]: + video = self.model(input) + return {'video': video} + + def postprocess(self, inputs: Dict[str, Any], **post_params) -> Dict[str, Any]: + video = tensor2vid(inputs['video']) + output_video_path = post_params.get('output_video', None) + temp_video_file = False + if output_video_path is None: + output_video_path = tempfile.NamedTemporaryFile(suffix='.mp4').name + temp_video_file = True + + fourcc = cv2.VideoWriter_fourcc(*'mp4v') + h, w, c = video[0].shape + video_writer = cv2.VideoWriter( + output_video_path, fourcc, fps=8, frameSize=(w, h)) + for i in range(len(video)): + img = cv2.cvtColor(video[i], cv2.COLOR_RGB2BGR) + video_writer.write(img) + video_writer.release() + if temp_video_file: + video_file_content = b'' + with open(output_video_path, 'rb') as f: + video_file_content = f.read() + os.remove(output_video_path) + return {OutputKeys.OUTPUT_VIDEO: video_file_content} + else: + return {OutputKeys.OUTPUT_VIDEO: output_video_path} + + +def tensor2vid(video, mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]): + mean = torch.tensor(mean, device=video.device).reshape(1, -1, 1, 1, 1) # ncfhw + std = torch.tensor(std, device=video.device).reshape(1, -1, 1, 1, 1) # ncfhw + + video = video.mul_(std).add_(mean) # unnormalize back to [0,1] + video.clamp_(0, 1) + video = video * 255.0 + + images = rearrange(video, 'b c f h w -> b f h w c')[0] + images = [(img.numpy()).astype('uint8') for img in images] + + return images diff --git a/modelscope/utils/constant.py b/modelscope/utils/constant.py index 0236d6c4..6cfc273c 100644 --- a/modelscope/utils/constant.py +++ b/modelscope/utils/constant.py @@ -256,6 +256,7 @@ class MultiModalTasks(object): text_to_video_synthesis = 'text-to-video-synthesis' efficient_diffusion_tuning = 'efficient-diffusion-tuning' multimodal_dialogue = 'multimodal-dialogue' + image_to_video_task = 'image-to-video-task' class ScienceTasks(object): diff --git a/requirements.txt b/requirements.txt index f5cec9e0..0832e6ab 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,2 +1 @@ -# test -r requirements/framework.txt diff --git a/test_image2video.py b/test_image2video.py new file mode 100644 index 00000000..b8842db5 --- /dev/null +++ b/test_image2video.py @@ -0,0 +1,38 @@ +import sys +import unittest + +from modelscope.models import Model +from modelscope.outputs import OutputKeys +from modelscope.pipelines import pipeline +from modelscope.utils.constant import DownloadMode, Tasks +from modelscope.utils.test_utils import test_level + + +class VideoDeinterlaceTest(unittest.TestCase): + + def setUp(self) -> None: + self.task = Tasks.image_to_video_task + self.model_id = '/mnt/workspace/Video_Generation/creative_space_image_to_video/models' + # self.model_revision = 'v1.0.1' + # self.dataset_id = 'buptwq/videocomposer-depths-style' + self.path = '/mnt/workspace/Video_Generation/creative_space_image_to_video/test.jpeg' + + @unittest.skipUnless(test_level() >= 1, 'skip test in current test level') + def test_run_pipeline(self): + pipe = pipeline(task=self.task, model=self.model_id) + # ds = MsDataset.load( + # self.dataset_id, + # split='train', + # download_mode=DownloadMode.FORCE_REDOWNLOAD) + # inputs = next(iter(ds)) + # inputs.update({'text': self.text}) + + inputs = {'img_path': self.path} + # _ = pipe(inputs) + + output_video_path = pipe(inputs, output_video='./output.mp4')[OutputKeys.OUTPUT_VIDEO] + print(output_video_path) + + +if __name__ == '__main__': + unittest.main() \ No newline at end of file diff --git a/tests/pipelines/test_image2video.py b/tests/pipelines/test_image2video.py new file mode 100644 index 00000000..b8842db5 --- /dev/null +++ b/tests/pipelines/test_image2video.py @@ -0,0 +1,38 @@ +import sys +import unittest + +from modelscope.models import Model +from modelscope.outputs import OutputKeys +from modelscope.pipelines import pipeline +from modelscope.utils.constant import DownloadMode, Tasks +from modelscope.utils.test_utils import test_level + + +class VideoDeinterlaceTest(unittest.TestCase): + + def setUp(self) -> None: + self.task = Tasks.image_to_video_task + self.model_id = '/mnt/workspace/Video_Generation/creative_space_image_to_video/models' + # self.model_revision = 'v1.0.1' + # self.dataset_id = 'buptwq/videocomposer-depths-style' + self.path = '/mnt/workspace/Video_Generation/creative_space_image_to_video/test.jpeg' + + @unittest.skipUnless(test_level() >= 1, 'skip test in current test level') + def test_run_pipeline(self): + pipe = pipeline(task=self.task, model=self.model_id) + # ds = MsDataset.load( + # self.dataset_id, + # split='train', + # download_mode=DownloadMode.FORCE_REDOWNLOAD) + # inputs = next(iter(ds)) + # inputs.update({'text': self.text}) + + inputs = {'img_path': self.path} + # _ = pipe(inputs) + + output_video_path = pipe(inputs, output_video='./output.mp4')[OutputKeys.OUTPUT_VIDEO] + print(output_video_path) + + +if __name__ == '__main__': + unittest.main() \ No newline at end of file