upgrade from 1.2.0 to 1.3.0

This commit is contained in:
maochaojie
2024-11-19 19:20:02 +08:00
parent 4ec0492897
commit a683061c6f
87 changed files with 10400 additions and 1117 deletions
@@ -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)
+258 -5
View File
@@ -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)
+167 -2
View File
@@ -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]