update 1.4.0
This commit is contained in:
@@ -1,7 +1,27 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from typing import TYPE_CHECKING
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
|
||||
from .diffusions import BaseDiffusion, DiffusionFluxRF
|
||||
from .samplers import BaseDiffusionSampler, DDIMSampler, FlowEluerSampler
|
||||
from .schedules import (BaseNoiseScheduler, FlowMatchShiftScheduler,
|
||||
ScaledLinearScheduler)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .diffusions import BaseDiffusion, DiffusionFluxRF
|
||||
from .samplers import BaseDiffusionSampler, DDIMSampler, FlowEluerSampler
|
||||
from .schedules import (BaseNoiseScheduler, FlowMatchShiftScheduler,
|
||||
ScaledLinearScheduler)
|
||||
else:
|
||||
_import_structure = {
|
||||
'diffusions': ['BaseDiffusion', 'DiffusionFluxRF'],
|
||||
'samplers': ['BaseDiffusionSampler', 'DDIMSampler', 'FlowEluerSampler'],
|
||||
'schedules': ['BaseNoiseScheduler', 'FlowMatchShiftScheduler',
|
||||
'ScaledLinearScheduler']
|
||||
}
|
||||
|
||||
import sys
|
||||
sys.modules[__name__] = LazyImportModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
|
||||
@@ -34,6 +34,7 @@ class BaseDiffusion(object):
|
||||
|
||||
def init_params(self):
|
||||
self.prediction_type = self.cfg.get('PREDICTION_TYPE', 'eps')
|
||||
self.use_dynamic_cfg = self.cfg.get('USE_DYNAMIC_CFG', False)
|
||||
self.noise_scheduler = NOISE_SCHEDULERS.build(self.cfg.NOISE_SCHEDULER,
|
||||
logger=self.logger)
|
||||
self.sampler_scheduler = NOISE_SCHEDULERS.build(self.cfg.get(
|
||||
@@ -56,7 +57,6 @@ class BaseDiffusion(object):
|
||||
model_kwargs={},
|
||||
steps=20,
|
||||
sampler=None,
|
||||
use_dynamic_cfg=False,
|
||||
guide_scale=None,
|
||||
guide_rescale=None,
|
||||
show_progress=False,
|
||||
@@ -79,7 +79,7 @@ class BaseDiffusion(object):
|
||||
if guide_scale is None or guide_scale == 1.0:
|
||||
out = model(x=x_t, t=t, **model_kwargs)
|
||||
else:
|
||||
if use_dynamic_cfg:
|
||||
if self.use_dynamic_cfg:
|
||||
guidance_scale = 1 + guide_scale * (
|
||||
(1 - math.cos(math.pi * (
|
||||
(steps - timestamp.item()) / steps)**5.0)) / 2)
|
||||
@@ -158,14 +158,16 @@ class BaseDiffusion(object):
|
||||
|
||||
def get_sampler(self, sampler):
|
||||
if isinstance(sampler, str):
|
||||
if sampler not in DIFFUSION_SAMPLERS.class_map:
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
if (not LazyImportModule.get_module_type(('DIFFUSION_SAMPLERS', sampler))) and (
|
||||
sampler not in DIFFUSION_SAMPLERS.class_map):
|
||||
if self.logger is not None:
|
||||
self.logger.info(
|
||||
f'{sampler} not in the defined samplers list {DIFFUSION_SAMPLERS.class_map.keys()}'
|
||||
f'{sampler} not in the defined samplers list.'
|
||||
)
|
||||
else:
|
||||
print(
|
||||
f'{sampler} not in the defined samplers list {DIFFUSION_SAMPLERS.class_map.keys()}'
|
||||
f'{sampler} not in the defined samplers list.'
|
||||
)
|
||||
return None
|
||||
sampler_cfg = Config(cfg_dict={'NAME': sampler}, load=False)
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import math
|
||||
import random
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable
|
||||
|
||||
@@ -30,7 +31,7 @@ class ScheduleOutput(object):
|
||||
|
||||
@NOISE_SCHEDULERS.register_class()
|
||||
class BaseNoiseScheduler(object):
|
||||
'''
|
||||
r'''
|
||||
In the diffusion model, the parameters related to the noise schedule are alpha, beta,
|
||||
and sigma. The following are the definitions of the above three parameters, which should
|
||||
be the basic property for the instance of noise scheduler.
|
||||
@@ -483,6 +484,14 @@ class FlowMatchFluxShiftScheduler(FlowMatchUniformScheduler):
|
||||
'MAX_SHIFT': {
|
||||
'value': 1.15,
|
||||
'description': 'The max shift factor for the timestamp.'
|
||||
},
|
||||
'PRE_T_SAMPLE': {
|
||||
'value': False,
|
||||
'description': 'Use pre-sampled timesteps or not, default is False.'
|
||||
},
|
||||
'PRE_T_SAMPLE_FOLD': {
|
||||
'value': 1,
|
||||
'description': 'The folds of pre-sampled timesteps.'
|
||||
}
|
||||
}
|
||||
|
||||
@@ -492,6 +501,23 @@ class FlowMatchFluxShiftScheduler(FlowMatchUniformScheduler):
|
||||
self.sigmoid_scale = self.cfg.get('SIGMOID_SCALE', 1)
|
||||
self.base_shift = self.cfg.get('BASE_SHIFT', 0.5)
|
||||
self.max_shift = self.cfg.get('MAX_SHIFT', 1.15)
|
||||
self.pre_t_sample = self.cfg.get('PRE_T_SAMPLE', False)
|
||||
self.pre_t_sample_fold = self.cfg.get('PRE_T_SAMPLE_FOLD', 1)
|
||||
if self.pre_t_sample:
|
||||
t = torch.sigmoid(torch.randn((self.num_timesteps * self.pre_t_sample_fold,)))
|
||||
# Scale and reverse the values to go from 1000 to 0
|
||||
timesteps = ((1 - t) * 1000)
|
||||
# Sort the timesteps in descending order
|
||||
self.pre_sample_timesteps, _ = torch.sort(timesteps, descending=True)
|
||||
else:
|
||||
self.pre_sample_timesteps = None
|
||||
|
||||
@property
|
||||
def pre_timesteps(self):
|
||||
fold_id = random.randint(0, self.pre_t_sample_fold - 1)
|
||||
# print("fold_id", fold_id)
|
||||
return self.pre_sample_timesteps[fold_id::self.pre_t_sample_fold]
|
||||
|
||||
|
||||
def time_shift(self, mu: float, sigma_scale: float, t: Tensor):
|
||||
return math.exp(mu) / (math.exp(mu) + (1 / t - 1)**sigma_scale)
|
||||
@@ -516,11 +542,22 @@ class FlowMatchFluxShiftScheduler(FlowMatchUniformScheduler):
|
||||
n, _, h, w = x_0.shape
|
||||
seq_len = (h // 2 * w // 2)
|
||||
if t is None:
|
||||
logits_norm = torch.randn(x_0.shape[0], device=x_0.device)
|
||||
logits_norm = logits_norm * self.sigmoid_scale # larger scale for more uniform sampling
|
||||
t = logits_norm.sigmoid() * self.num_timesteps
|
||||
if self.pre_t_sample:
|
||||
timestep_indices = torch.randint(
|
||||
1,
|
||||
self.num_timesteps - 1,
|
||||
(x_0.shape[0],)
|
||||
)
|
||||
timestep_indices = timestep_indices.long()
|
||||
t = [self.pre_timesteps[x.item()].to(x_0.device) for x in timestep_indices]
|
||||
t = torch.stack(t, dim=0)
|
||||
else:
|
||||
logits_norm = torch.randn(x_0.shape[0], device=x_0.device)
|
||||
logits_norm = logits_norm * self.sigmoid_scale # larger scale for more uniform sampling
|
||||
t = logits_norm.sigmoid() * self.num_timesteps
|
||||
sigma = self.t_to_sigma(t, seq_len=seq_len)
|
||||
shape = (x_0.size(0), ) + (1, ) * (x_0.ndim - 1)
|
||||
# print(sigma)
|
||||
x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise
|
||||
return ScheduleOutput(x_0=x_0,
|
||||
x_t=x_t,
|
||||
|
||||
Reference in New Issue
Block a user