upgrade from 1.2.0 to 1.3.0
This commit is contained in:
@@ -1,3 +1,4 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.network.autoencoder.ae_kl import AutoencoderKL
|
||||
from scepter.modules.model.network.autoencoder.ae_kl_cogvideox import AutoencoderKLCogVideoX
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,7 +1,8 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.network.ldm.ldm import LatentDiffusion
|
||||
from scepter.modules.model.network.ldm.ldm_ace import LatentDiffusionACE
|
||||
from scepter.modules.model.network.ldm.ldm_ace import (LatentDiffusionACE,
|
||||
LatentDiffusionACERefiner)
|
||||
from scepter.modules.model.network.ldm.ldm_edit import LatentDiffusionEdit
|
||||
from scepter.modules.model.network.ldm.ldm_pixart import LatentDiffusionPixart
|
||||
from scepter.modules.model.network.ldm.ldm_sce import (
|
||||
@@ -9,3 +10,6 @@ from scepter.modules.model.network.ldm.ldm_sce import (
|
||||
LatentDiffusionXLSCEControl, LatentDiffusionXLSCETuning)
|
||||
from scepter.modules.model.network.ldm.ldm_sd3 import LatentDiffusionSD3
|
||||
from scepter.modules.model.network.ldm.ldm_xl import LatentDiffusionXL
|
||||
from scepter.modules.model.network.ldm.ldm_cogvideox import LatentDiffusionCogVideoX
|
||||
from scepter.modules.model.network.ldm.ldm_flux import (LatentDiffusionFlux,
|
||||
LatentDiffusionFluxMR)
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import math
|
||||
import random
|
||||
from contextlib import nullcontext
|
||||
|
||||
@@ -10,6 +11,7 @@ from torch import nn
|
||||
|
||||
from scepter.modules.model.network.ldm import LatentDiffusion
|
||||
from scepter.modules.model.registry import MODELS
|
||||
import torchvision.transforms as T
|
||||
from scepter.modules.model.utils.basic_utils import check_list_of_list
|
||||
from scepter.modules.model.utils.basic_utils import \
|
||||
pack_imagelist_into_tensor_v2 as pack_imagelist_into_tensor
|
||||
@@ -67,10 +69,10 @@ class LatentDiffusionACE(LatentDiffusion):
|
||||
if self.use_text_pos_embeddings and not torch.sum(
|
||||
self.text_position_embeddings.pos) > 0:
|
||||
identifier_cont, identifier_cont_mask = getattr(
|
||||
self.cond_stage_model, 'encode')(self.text_indentifers,
|
||||
self.cond_stage_model, 'encode_list_of_list')(self.text_indentifers,
|
||||
return_mask=True)
|
||||
self.text_position_embeddings.load_state_dict(
|
||||
{'pos': identifier_cont[:, 0, :]})
|
||||
{'pos': torch.cat( [one_id[0][0, :].unsqueeze(0) for one_id in identifier_cont], dim=0)})
|
||||
cont_, cont_mask_ = [], []
|
||||
for pp, edit, c, cm in zip(prompt, edit_image, cont, cont_mask):
|
||||
if isinstance(pp, list):
|
||||
@@ -138,7 +140,7 @@ class LatentDiffusionACE(LatentDiffusion):
|
||||
prompt_ = [[pp] if isinstance(pp, str) else pp for pp in prompt]
|
||||
try:
|
||||
cont, cont_mask = getattr(self.cond_stage_model,
|
||||
'encode_list')(prompt_, return_mask=True)
|
||||
'encode_list_of_list')(prompt_, return_mask=True)
|
||||
except Exception as e:
|
||||
print(e, prompt_)
|
||||
cont, cont_mask = self.cond_stage_embeddings(prompt, edit_image, cont,
|
||||
@@ -240,11 +242,11 @@ class LatentDiffusionACE(LatentDiffusion):
|
||||
# with torch.autocast(device_type="cuda", enabled=True, dtype=torch.bfloat16):
|
||||
prompt_ = [[pp] if isinstance(pp, str) else pp for pp in prompt]
|
||||
cont, cont_mask = getattr(self.cond_stage_model,
|
||||
'encode_list')(prompt_, return_mask=True)
|
||||
'encode_list_of_list')(prompt_, return_mask=True)
|
||||
cont, cont_mask = self.cond_stage_embeddings(prompt, edit_image, cont,
|
||||
cont_mask)
|
||||
null_cont, null_cont_mask = getattr(self.cond_stage_model,
|
||||
'encode_list')(n_prompt,
|
||||
'encode_list_of_list')(n_prompt,
|
||||
return_mask=True)
|
||||
null_cont, null_cont_mask = self.cond_stage_embeddings(
|
||||
prompt, edit_image, null_cont, null_cont_mask)
|
||||
@@ -349,3 +351,254 @@ class LatentDiffusionACE(LatentDiffusion):
|
||||
__class__.__name__,
|
||||
LatentDiffusionACE.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@MODELS.register_class()
|
||||
class LatentDiffusionACERefiner(LatentDiffusionACE):
|
||||
def init_params(self):
|
||||
super().init_params()
|
||||
self.enhence_model_cfg = self.cfg.get("ENHENCE_MODEL", None)
|
||||
self.enhence_sampler_cfg = self.cfg.get("ENHENCE_SAMPLER_CFG", {})
|
||||
def construct_network(self):
|
||||
super().construct_network()
|
||||
if self.enhence_model_cfg:
|
||||
self.enhence_model = MODELS.build(self.enhence_model_cfg, logger=self.logger).eval().requires_grad_(False)
|
||||
self.enhence_sampler_cfg = {key.lower(): value for key, value in self.enhence_sampler_cfg.items()}
|
||||
else:
|
||||
self.enhence_model = None
|
||||
self.enhence_sampler_cfg = None
|
||||
|
||||
def forward_sample(self,
|
||||
edit_image=[],
|
||||
edit_mask=[],
|
||||
noise=None,
|
||||
cond_mask=[],
|
||||
x_shapes=[],
|
||||
prompt=[],
|
||||
n_prompt=[],
|
||||
sampler='ddim',
|
||||
sample_steps=20,
|
||||
seed=2023,
|
||||
guide_scale=4.5,
|
||||
guide_rescale=0.5,
|
||||
discretization='trailing',
|
||||
**kwargs
|
||||
):
|
||||
'''
|
||||
Args:
|
||||
edit_image: list of list of edit_image
|
||||
edit_image_mask: list of list of edit_image_mask
|
||||
image: target image
|
||||
image_mask: target image mask
|
||||
prompt: list of list of text
|
||||
n_prompt: list of list of text
|
||||
sampler:
|
||||
sample_steps:
|
||||
seed:
|
||||
guide_scale:
|
||||
guide_rescale:
|
||||
discretization:
|
||||
log_num:
|
||||
**kwargs:
|
||||
|
||||
Returns:
|
||||
|
||||
'''
|
||||
|
||||
# prepare data
|
||||
context, null_context = {}, {}
|
||||
context['x_shapes'] = null_context['x_shapes'] = x_shapes
|
||||
# process image mask
|
||||
|
||||
context['x_mask'] = null_context['x_mask'] = cond_mask
|
||||
# process text
|
||||
# with torch.autocast(device_type="cuda", enabled=True, dtype=torch.bfloat16):
|
||||
|
||||
cont, cont_mask = getattr(self.cond_stage_model, 'encode_list')(prompt, return_mask=True)
|
||||
cont, cont_mask = self.cond_stage_embeddings(prompt, edit_image, cont, cont_mask)
|
||||
null_cont, null_cont_mask = getattr(self.cond_stage_model, 'encode_list')(n_prompt, return_mask=True)
|
||||
null_cont, null_cont_mask = self.cond_stage_embeddings(prompt, edit_image, null_cont, null_cont_mask)
|
||||
context['crossattn'] = cont
|
||||
null_context['crossattn'] = null_cont
|
||||
|
||||
|
||||
null_context['edit'] = context['edit'] = edit_image
|
||||
null_context['edit_mask'] = context['edit_mask'] = edit_mask
|
||||
|
||||
# process sample
|
||||
model = self.model_ema if self.use_ema and self.eval_ema else self.model
|
||||
embedding_context = model.no_sync if isinstance(model, torch.distributed.fsdp.FullyShardedDataParallel) \
|
||||
else nullcontext
|
||||
with embedding_context():
|
||||
samples = self.diffusion.sample(solver=sampler,
|
||||
noise=noise,
|
||||
model=model,
|
||||
model_kwargs=[{
|
||||
'cond': context,
|
||||
'mask': cont_mask,
|
||||
'text_position_embeddings': self.text_position_embeddings.pos if hasattr(
|
||||
self.text_position_embeddings, 'pos') else None
|
||||
}, {
|
||||
'cond': null_context,
|
||||
'mask': null_cont_mask,
|
||||
'text_position_embeddings': self.text_position_embeddings.pos if hasattr(
|
||||
self.text_position_embeddings, 'pos') else None
|
||||
}] if guide_scale is not None and guide_scale > 1 else {
|
||||
'cond': context,
|
||||
'mask': cont_mask,
|
||||
'text_position_embeddings': self.text_position_embeddings.pos if hasattr(
|
||||
self.text_position_embeddings, 'pos') else None
|
||||
},
|
||||
cat_uc=False,
|
||||
steps=sample_steps,
|
||||
guide_scale=guide_scale,
|
||||
guide_rescale=guide_rescale,
|
||||
discretization=discretization,
|
||||
show_progress=True,
|
||||
seed=seed,
|
||||
condition_fn=None,
|
||||
clamp=None,
|
||||
percentile=None,
|
||||
t_max=None,
|
||||
t_min=None,
|
||||
discard_penultimate_step=None,
|
||||
return_intermediate=None,
|
||||
**kwargs)
|
||||
|
||||
samples = unpack_tensor_into_imagelist(samples, x_shapes)
|
||||
x_samples = self.decode_first_stage(samples)
|
||||
return x_samples
|
||||
|
||||
def upscale_resize(self, image, interpolation=T.InterpolationMode.BILINEAR):
|
||||
_, c, H, W = image.shape
|
||||
scale = max(1.0, math.sqrt(4096 / ((H / 16) * (W / 16))))
|
||||
rH = int(H * scale) // 16 * 16 # ensure divisible by self.d
|
||||
rW = int(W * scale) // 16 * 16
|
||||
image = T.Resize((rH, rW), interpolation=interpolation, antialias=True)(image)
|
||||
return image
|
||||
|
||||
@torch.no_grad()
|
||||
def forward_test(self,
|
||||
edit_image=[],
|
||||
edit_image_mask=[],
|
||||
image=None,
|
||||
image_mask=None,
|
||||
prompt=[],
|
||||
n_prompt=[],
|
||||
sampler='ddim',
|
||||
sample_steps=20,
|
||||
seed=2023,
|
||||
guide_scale=4.5,
|
||||
guide_rescale=0.5,
|
||||
discretization='trailing',
|
||||
enhance_scale=0.99,
|
||||
log_num=-1,
|
||||
**kwargs):
|
||||
assert check_list_of_list(prompt) and check_list_of_list(edit_image) and check_list_of_list(edit_image_mask)
|
||||
assert len(edit_image) == len(edit_image_mask) == len(prompt)
|
||||
assert self.cond_stage_model is not None
|
||||
# gc_seg is unused
|
||||
kwargs.pop("gc_seg", -1)
|
||||
prompt, n_prompt, image, image_mask, edit_image, edit_image_mask = self.limit_batch_data(
|
||||
[prompt, n_prompt, image, image_mask, edit_image, edit_image_mask], log_num)
|
||||
|
||||
prompt = [[pp] if isinstance(pp, str) else pp for pp in prompt]
|
||||
|
||||
g = torch.Generator(device=we.device_id)
|
||||
seed = seed if seed >= 0 else random.randint(0, 2 ** 32 - 1)
|
||||
g.manual_seed(seed)
|
||||
n_prompt = copy.deepcopy(prompt)
|
||||
# only modify the last prompt to be zero
|
||||
for nn_p_id, nn_p in enumerate(n_prompt):
|
||||
if isinstance(nn_p, str):
|
||||
n_prompt[nn_p_id] = [""]
|
||||
elif isinstance(nn_p, list):
|
||||
n_prompt[nn_p_id][-1] = ""
|
||||
else:
|
||||
raise NotImplementedError
|
||||
# process image
|
||||
image = to_device(image)
|
||||
x = self.encode_first_stage(image, **kwargs)
|
||||
noise = [torch.empty(*i.shape, device=we.device_id).normal_(generator=g) for i in x]
|
||||
noise, x_shapes = pack_imagelist_into_tensor(noise)
|
||||
image_mask = to_device(image_mask, strict=False)
|
||||
cond_mask = [self.interpolate_func(i) for i in image_mask] if image_mask is not None else [None] * len(image)
|
||||
|
||||
# processe edit image & edit image mask
|
||||
edit_image = [to_device(i, strict=False) for i in edit_image]
|
||||
edit_image_mask = [to_device(i, strict=False) for i in edit_image_mask]
|
||||
e_img, e_mask = [], []
|
||||
for u, m in zip(edit_image, edit_image_mask):
|
||||
if u is None:
|
||||
continue
|
||||
if m is None:
|
||||
m = [None] * len(u)
|
||||
e_img.append(self.encode_first_stage(u, **kwargs))
|
||||
e_mask.append([self.interpolate_func(i) for i in m])
|
||||
|
||||
x_samples = self.forward_sample(
|
||||
edit_image=e_img,
|
||||
edit_mask=e_mask,
|
||||
noise=noise,
|
||||
cond_mask=cond_mask,
|
||||
x_shapes=x_shapes,
|
||||
prompt=prompt,
|
||||
n_prompt=n_prompt,
|
||||
sampler=sampler,
|
||||
sample_steps=sample_steps,
|
||||
seed=seed,
|
||||
guide_scale=guide_scale,
|
||||
guide_rescale=guide_rescale,
|
||||
discretization='trailing',
|
||||
**kwargs)
|
||||
|
||||
if self.enhence_model and enhance_scale > 0:
|
||||
x_samples = [self.upscale_resize(x) for x in x_samples]
|
||||
x_start = self.enhence_model.encode_first_stage(x_samples, **kwargs)
|
||||
noise = []
|
||||
for i, x in enumerate(x_start):
|
||||
noise_ = self.enhence_model.noise_sample(1, x_samples[i].shape[2], x_samples[i].shape[3], seed)
|
||||
noise.append(noise_)
|
||||
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
|
||||
x_samples = self.enhence_model.forward_sample(noise = noise,
|
||||
x = x_start,
|
||||
reverse_scale = enhance_scale,
|
||||
prompt =[kwargs.pop("enhance_prompt", "") for _ in noise],
|
||||
**self.enhence_sampler_cfg)
|
||||
outputs = list()
|
||||
for i in range(len(prompt)):
|
||||
rec_img = torch.clamp((x_samples[i].float() + 1.0) / 2.0 + self.decoder_bias / 255, min=0.0, max=1.0)
|
||||
rec_img = rec_img.squeeze(0)
|
||||
edit_imgs, edit_img_masks = [], []
|
||||
if edit_image is not None and edit_image[i] is not None:
|
||||
if edit_image_mask[i] is None:
|
||||
edit_image_mask[i] = [None] * len(edit_image[i])
|
||||
for edit_img, edit_mask in zip(edit_image[i], edit_image_mask[i]):
|
||||
edit_img = torch.clamp((edit_img + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
edit_imgs.append(edit_img.squeeze(0))
|
||||
if edit_mask is None:
|
||||
edit_mask = torch.ones_like(edit_img[[0], :, :])
|
||||
edit_img_masks.append(edit_mask)
|
||||
one_tup = {
|
||||
'reconstruct_image': rec_img,
|
||||
'instruction': prompt[i],
|
||||
'edit_image': edit_imgs if len(edit_imgs) > 0 else None,
|
||||
'edit_mask': edit_img_masks if len(edit_imgs) > 0 else None
|
||||
}
|
||||
if image is not None:
|
||||
if image_mask is None:
|
||||
image_mask = [None] * len(image)
|
||||
ori_img = torch.clamp((image[i] + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
one_tup['target_image'] = ori_img.squeeze(0)
|
||||
one_tup['target_mask'] = image_mask[i] if image_mask[i] is not None else torch.ones_like(
|
||||
ori_img[[0], :, :])
|
||||
outputs.append(one_tup)
|
||||
|
||||
return outputs
|
||||
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('MODEL',
|
||||
__class__.__name__,
|
||||
LatentDiffusionACERefiner.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@@ -0,0 +1,225 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import random
|
||||
import torch
|
||||
from typing import Tuple
|
||||
|
||||
from scepter.modules.model.network.ldm import LatentDiffusion
|
||||
from scepter.modules.model.registry import MODELS
|
||||
from scepter.modules.model.utils.basic_utils import default
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.model.backbone.cogvideox.utils import get_3d_rotary_pos_embed, get_resize_crop_region_for_grid
|
||||
|
||||
|
||||
@MODELS.register_class()
|
||||
class LatentDiffusionCogVideoX(LatentDiffusion):
|
||||
para_dict = LatentDiffusion.para_dict
|
||||
|
||||
def init_params(self):
|
||||
super().init_params()
|
||||
self.latent_channels = self.model_config.get('LATENT_CHANNELS', self.model_config.IN_CHANNELS)
|
||||
self.scale_factor_spatial = self.cfg.get('SCALE_FACTOR_SPATIAL', 8)
|
||||
self.scale_factor_temporal = self.cfg.get('SCALE_FACTOR_TEMPORAL', 4)
|
||||
self.scaling_factor_image = self.cfg.get('SCALING_FACTOR_IMAGE', 0.7)
|
||||
self.use_rotary_positional_embeddings = self.model_config.get('USE_ROTARY_POSITIONAL_EMBEDDINGS', False)
|
||||
self.attention_head_dim = self.model_config.get('ATTENTION_HEAD_DIM', 64)
|
||||
self.patch_size = self.model_config.get('PATCH_SIZE', 2)
|
||||
self.sample_height = self.first_stage_config.get('SAMPLE_HEIGHT', 480)
|
||||
self.sample_width = self.first_stage_config.get('SAMPLE_WIDTH', 720)
|
||||
self.noised_image_dropout = self.cfg.get('NOISED_IMAGE_DROPOUT', 0.05)
|
||||
|
||||
def construct_network(self):
|
||||
super().construct_network()
|
||||
self.model = self.model.to(getattr(torch, self.model_config.DTYPE))
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_first_stage(self, x, **kwargs):
|
||||
if isinstance(x, list):
|
||||
x = torch.stack(x, dim=0) # [B, C, F, H, W]
|
||||
latents = self.scaling_factor_image * self.first_stage_model.encode(x).sample()
|
||||
return latents
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_first_stage(self, latents):
|
||||
latents = latents.permute(0, 2, 1, 3, 4) # [batch_size, num_channels, num_frames, height, width]
|
||||
latents = 1 / self.scaling_factor_image * latents
|
||||
frames = self.first_stage_model.decode(latents)
|
||||
return frames
|
||||
|
||||
def get_image_latent(self, image, video, noise):
|
||||
latent = torch.zeros_like(noise)
|
||||
if isinstance(image, list):
|
||||
image = torch.stack(image, dim=0) # [B, C, F, H, W]
|
||||
if len(image.shape) == 4: # [B, C, H, W]
|
||||
image = image.unsqueeze(2) # [B, C, F, H, W]
|
||||
image_latent = self.encode_first_stage(image) # [B, C, F, H, W]
|
||||
image_latent = image_latent.permute(0, 2, 1, 3, 4) # [B, F, C, H, W]
|
||||
latent[:, :1, :, :, :] = image_latent
|
||||
return latent, image
|
||||
|
||||
def noise_sample(self, batch_size, num_frames, height, width, generator, dtype=torch.bfloat16):
|
||||
shape = (batch_size,
|
||||
(num_frames - 1) // self.scale_factor_temporal + 1,
|
||||
self.latent_channels,
|
||||
height // self.scale_factor_spatial,
|
||||
width // self.scale_factor_spatial
|
||||
)
|
||||
noise = torch.randn(shape, generator=generator, dtype=dtype, device='cpu').to(we.device_id)
|
||||
return noise
|
||||
|
||||
def _prepare_rotary_positional_embeddings(
|
||||
self,
|
||||
height: int,
|
||||
width: int,
|
||||
num_frames: int,
|
||||
device: torch.device,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
grid_height = height // (self.scale_factor_spatial * self.patch_size)
|
||||
grid_width = width // (self.scale_factor_spatial * self.patch_size)
|
||||
base_size_width = self.sample_width // (self.scale_factor_spatial * self.patch_size)
|
||||
base_size_height = self.sample_height // (self.scale_factor_spatial * self.patch_size)
|
||||
|
||||
grid_crops_coords = get_resize_crop_region_for_grid(
|
||||
(grid_height, grid_width), base_size_width, base_size_height
|
||||
)
|
||||
freqs_cos, freqs_sin = get_3d_rotary_pos_embed(
|
||||
embed_dim=self.attention_head_dim,
|
||||
crops_coords=grid_crops_coords,
|
||||
grid_size=(grid_height, grid_width),
|
||||
temporal_size=num_frames,
|
||||
)
|
||||
|
||||
freqs_cos = freqs_cos.to(device=device)
|
||||
freqs_sin = freqs_sin.to(device=device)
|
||||
return freqs_cos, freqs_sin
|
||||
|
||||
|
||||
def forward_train(self, video=None, video_latent=None, image=None, noise=None, prompt=None, image_size=None, **kwargs):
|
||||
# video: [B, C, F, H, W]
|
||||
if image_size is None: image_size = [480, 720]
|
||||
if video_latent is not None:
|
||||
x_start = torch.stack(video_latent)
|
||||
else:
|
||||
x_start = self.encode_first_stage(video, **kwargs)
|
||||
x_start = x_start.permute(0, 2, 1, 3, 4) # [B, F, C, H, W]
|
||||
t = torch.randint(low=0, high=self.num_timesteps, size=(len(video),), device=we.device_id)
|
||||
|
||||
if prompt and self.cond_stage_model:
|
||||
with torch.autocast(device_type='cuda', enabled=True, dtype=torch.bfloat16):
|
||||
cont = getattr(self.cond_stage_model, 'encode')(prompt, return_mask=False, use_mask=False)
|
||||
|
||||
if noise is None:
|
||||
noise = torch.randn_like(x_start)
|
||||
|
||||
if image is not None:
|
||||
if random.random() < self.noised_image_dropout:
|
||||
image_latent = torch.zeros_like(noise)
|
||||
else:
|
||||
image_latent, _ = self.get_image_latent(image, video, noise)
|
||||
else:
|
||||
image_latent = None
|
||||
|
||||
height, width = image_size
|
||||
image_rotary_emb = (
|
||||
self._prepare_rotary_positional_embeddings(height, width, noise.size(1), we.device_id)
|
||||
if self.use_rotary_positional_embeddings
|
||||
else None
|
||||
)
|
||||
|
||||
loss = self.diffusion.loss(x_0=x_start,
|
||||
t=t,
|
||||
model=self.model,
|
||||
model_kwargs={"cond": cont,
|
||||
'image_latent': image_latent,
|
||||
'image_rotary_emb': image_rotary_emb},
|
||||
noise=noise,
|
||||
**kwargs)
|
||||
loss = loss.mean()
|
||||
ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
|
||||
return ret
|
||||
|
||||
@torch.no_grad()
|
||||
@torch.autocast('cuda', dtype=torch.bfloat16)
|
||||
def forward_test(self,
|
||||
video=None,
|
||||
image=None,
|
||||
prompt=None,
|
||||
n_prompt=None,
|
||||
sampler='ddim',
|
||||
sample_steps=50,
|
||||
seed=42,
|
||||
guide_scale=6.0,
|
||||
guide_rescale=0.0,
|
||||
num_frames=49,
|
||||
image_size=None,
|
||||
show_process=False,
|
||||
**kwargs):
|
||||
if image_size is None:
|
||||
image_size = [480, 720]
|
||||
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
||||
generator = torch.Generator().manual_seed(seed)
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
num_samples = len(prompt)
|
||||
n_prompt = default(n_prompt, [self.default_n_prompt] * len(prompt))
|
||||
|
||||
if prompt and self.cond_stage_model:
|
||||
with torch.autocast(device_type='cuda', enabled=True, dtype=torch.bfloat16):
|
||||
cont = getattr(self.cond_stage_model, 'encode')(prompt, return_mask=False, use_mask=False)
|
||||
null_cont = getattr(self.cond_stage_model, 'encode')(n_prompt, return_mask=False, use_mask=False)
|
||||
|
||||
height, width = image_size
|
||||
noise = self.noise_sample(num_samples, num_frames, height, width, generator)
|
||||
image_rotary_emb = (
|
||||
self._prepare_rotary_positional_embeddings(height, width, noise.size(1), we.device_id)
|
||||
if self.use_rotary_positional_embeddings
|
||||
else None
|
||||
)
|
||||
image_latent, image = self.get_image_latent(image, video, noise) if image is not None else (None, None)
|
||||
|
||||
samples = self.diffusion.sample(noise=noise,
|
||||
sampler=sampler,
|
||||
model=self.model,
|
||||
model_kwargs=[{
|
||||
'cond': cont,
|
||||
'image_latent': image_latent,
|
||||
'image_rotary_emb': image_rotary_emb,
|
||||
}, {
|
||||
'cond': null_cont,
|
||||
'image_latent': image_latent,
|
||||
'image_rotary_emb': image_rotary_emb,
|
||||
}],
|
||||
steps=sample_steps,
|
||||
show_progress=True,
|
||||
use_dynamic_cfg=True,
|
||||
guide_scale=guide_scale,
|
||||
guide_rescale=guide_rescale,
|
||||
return_intermediate=None,
|
||||
**kwargs).float()
|
||||
|
||||
x_frames = self.decode_first_stage(samples).float()
|
||||
|
||||
outputs = []
|
||||
for batch_idx in range(num_samples):
|
||||
rec_video = torch.clamp(x_frames[batch_idx] / 2 + 0.5, min=0.0, max=1.0)
|
||||
one_tup = {
|
||||
'reconstruct_video': rec_video.squeeze(0).float(),
|
||||
'instruction': prompt[batch_idx]
|
||||
}
|
||||
if image is not None:
|
||||
ori_image = torch.clamp(image[batch_idx] / 2 + 0.5, min=0.0, max=1.0)
|
||||
one_tup['edit_image'] = ori_image
|
||||
if video is not None:
|
||||
ori_video = torch.clamp(video[batch_idx] / 2 + 0.5, min=0.0, max=1.0)
|
||||
one_tup['target_video'] = ori_video.squeeze(0)
|
||||
outputs.append(one_tup)
|
||||
return outputs
|
||||
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('MODEL',
|
||||
__class__.__name__,
|
||||
LatentDiffusionCogVideoX.para_dict,
|
||||
set_name=True)
|
||||
@@ -4,15 +4,19 @@ import copy
|
||||
import math
|
||||
import numbers
|
||||
import random
|
||||
from contextlib import nullcontext
|
||||
|
||||
import torch
|
||||
from scepter.modules.model.network.ldm import LatentDiffusion
|
||||
from scepter.modules.model.registry import MODELS, BACKBONES, LOSSES, TOKENIZERS, EMBEDDERS, DIFFUSIONS
|
||||
from scepter.modules.model.utils.basic_utils import disabled_train
|
||||
from scepter.modules.model.utils.basic_utils import disabled_train, check_list_of_list, to_device, \
|
||||
pack_imagelist_into_tensor, unpack_tensor_into_imagelist, limit_batch_data
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.model.utils.basic_utils import count_params
|
||||
|
||||
|
||||
|
||||
@MODELS.register_class()
|
||||
class LatentDiffusionFlux(LatentDiffusion):
|
||||
para_dict = LatentDiffusion.para_dict
|
||||
@@ -137,7 +141,7 @@ class LatentDiffusionFlux(LatentDiffusion):
|
||||
def forward_test(self,
|
||||
image=None,
|
||||
prompt=None,
|
||||
sampler='flow_eluer',
|
||||
sampler='flow_euler',
|
||||
sample_steps=20,
|
||||
seed=2023,
|
||||
guide_scale=4.5,
|
||||
@@ -218,3 +222,164 @@ class LatentDiffusionFlux(LatentDiffusion):
|
||||
@torch.no_grad()
|
||||
def decode_first_stage(self, z):
|
||||
return self.first_stage_model.decode(z)
|
||||
|
||||
@MODELS.register_class()
|
||||
class LatentDiffusionFluxMR(LatentDiffusionFlux):
|
||||
para_dict = {
|
||||
}
|
||||
para_dict.update(LatentDiffusion.para_dict)
|
||||
def forward_train(self,
|
||||
image=None,
|
||||
noise=None,
|
||||
prompt=[],
|
||||
**kwargs):
|
||||
if check_list_of_list(prompt):
|
||||
prompt = [pp[0] for pp in prompt]
|
||||
assert self.cond_stage_model is not None
|
||||
gc_seg = kwargs.pop("gc_seg", [])
|
||||
gc_seg = int(gc_seg[0]) if len(gc_seg) > 0 else 0
|
||||
context = getattr(self.cond_stage_model, 'encode')(prompt)
|
||||
|
||||
image = to_device(image)
|
||||
x_start = self.encode_first_stage(image, **kwargs)
|
||||
loss_mask, _ = pack_imagelist_into_tensor(tuple(torch.ones_like(ix, dtype=torch.bool, device=ix.device) for ix in x_start))
|
||||
x_start, x_shapes = pack_imagelist_into_tensor(x_start)
|
||||
context['x_shapes'] = x_shapes
|
||||
guide_scale = self.guide_scale
|
||||
if guide_scale is not None:
|
||||
guide_scale = torch.full((x_start.shape[0],), guide_scale, device=x_start.device, dtype=x_start.dtype)
|
||||
else:
|
||||
guide_scale = None
|
||||
loss = self.diffusion.loss(x_0=x_start,
|
||||
model=self.model,
|
||||
model_kwargs={"cond": context,
|
||||
"gc_seg": gc_seg,
|
||||
"guidance": guide_scale},
|
||||
noise=None,
|
||||
reduction='none',
|
||||
**kwargs)
|
||||
loss = loss[loss_mask].mean()
|
||||
ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
|
||||
return ret
|
||||
|
||||
@torch.no_grad()
|
||||
def forward_sample(self,
|
||||
noise = None,
|
||||
prompt=None,
|
||||
sampler='flow_euler',
|
||||
sample_steps=20,
|
||||
guide_scale=3.5,
|
||||
show_process=True,
|
||||
x = None,
|
||||
reverse_scale = 0.,
|
||||
**kwargs
|
||||
):
|
||||
noise, x_shapes = pack_imagelist_into_tensor(noise)
|
||||
if x is not None:
|
||||
x, _ = pack_imagelist_into_tensor(x)
|
||||
context = getattr(self.cond_stage_model, 'encode')(prompt)
|
||||
context["x_shapes"] = x_shapes
|
||||
guide_scale = guide_scale or self.guide_scale
|
||||
if guide_scale is not None:
|
||||
guide_scale = torch.full((noise.shape[0],), guide_scale, device=noise.device, dtype=noise.dtype)
|
||||
else:
|
||||
guide_scale = None
|
||||
# UNet use input n_prompt
|
||||
model = self.model_ema if self.use_ema and self.eval_ema else self.model
|
||||
embedding_context = model.no_sync if isinstance(model, torch.distributed.fsdp.FullyShardedDataParallel) \
|
||||
else nullcontext
|
||||
with embedding_context():
|
||||
x_samples = self.diffusion.sample(
|
||||
noise=noise,
|
||||
sampler=sampler,
|
||||
model=self.model,
|
||||
model_kwargs={"cond": context, "guidance": guide_scale, "gc_seg": -1},
|
||||
steps=sample_steps,
|
||||
show_progress=True,
|
||||
guide_scale=guide_scale,
|
||||
return_intermediate=None,
|
||||
reverse_scale = reverse_scale,
|
||||
x = x,
|
||||
**kwargs).float()
|
||||
x_samples = unpack_tensor_into_imagelist(x_samples, x_shapes)
|
||||
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
|
||||
x_samples = self.decode_first_stage(x_samples)
|
||||
return x_samples
|
||||
@torch.no_grad()
|
||||
def forward_test(self,
|
||||
image=None,
|
||||
prompt=[],
|
||||
sampler='flow_euler',
|
||||
sample_steps=20,
|
||||
seed=2023,
|
||||
guide_scale=3.5,
|
||||
guide_rescale=0.0,
|
||||
show_process=True,
|
||||
log_num = -1,
|
||||
**kwargs):
|
||||
|
||||
if check_list_of_list(prompt):
|
||||
prompt = [pp[0] for pp in prompt]
|
||||
assert self.cond_stage_model is not None
|
||||
# gc_seg is unused
|
||||
prompt, image = limit_batch_data([prompt, image], log_num)
|
||||
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
||||
|
||||
if 'index' in kwargs:
|
||||
kwargs.pop('index')
|
||||
if image is not None:
|
||||
noise = [self.noise_sample(1, ix.shape[1], ix.shape[2], seed) for ix in image]
|
||||
else:
|
||||
image_size = None
|
||||
if 'meta' in kwargs:
|
||||
meta = kwargs.pop('meta')
|
||||
if 'image_size' in meta:
|
||||
h = int(meta['image_size'][0][0])
|
||||
w = int(meta['image_size'][1][0])
|
||||
image_size = [h, w]
|
||||
if 'image_size' in kwargs:
|
||||
image_size = kwargs.pop('image_size')
|
||||
if isinstance(image_size, numbers.Number):
|
||||
image_size = [image_size, image_size]
|
||||
if image_size is None:
|
||||
image_size = [1024, 1024]
|
||||
height, width = image_size
|
||||
noise = [self.noise_sample(1, height, width, seed) for _ in prompt]
|
||||
|
||||
x_samples = self.forward_sample(
|
||||
prompt=prompt,
|
||||
sampler=sampler,
|
||||
sample_steps=sample_steps,
|
||||
guide_scale=guide_scale,
|
||||
show_process=show_process,
|
||||
noise=noise,
|
||||
)
|
||||
|
||||
|
||||
outputs = list()
|
||||
for i in range(len(prompt)):
|
||||
rec_img = torch.clamp((x_samples[i].float() + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
rec_img = rec_img.squeeze(0)
|
||||
one_tup = {'prompt': prompt[i], 'n_prompt': '', 'image': rec_img}
|
||||
outputs.append(one_tup)
|
||||
return outputs
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('MODEL',
|
||||
__class__.__name__,
|
||||
LatentDiffusionFlux.para_dict,
|
||||
set_name=True)
|
||||
@torch.no_grad()
|
||||
def encode_first_stage(self, x, **kwargs):
|
||||
def run_one_image(u):
|
||||
zu = self.first_stage_model.encode(u)
|
||||
if isinstance(zu, (tuple, list)):
|
||||
zu = zu[0]
|
||||
return zu
|
||||
|
||||
z = [run_one_image(u.unsqueeze(0) if u.dim == 3 else u) for u in x]
|
||||
return z
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_first_stage(self, z):
|
||||
return [self.first_stage_model.decode(zu) for zu in z]
|
||||
Reference in New Issue
Block a user