update
This commit is contained in:
@@ -46,19 +46,23 @@ class BaseDiffusion(object):
|
||||
|
||||
|
||||
def sample(self, noise, model, model_kwargs={}, steps=20, sampler=None, use_dynamic_cfg=False, guide_scale=None, guide_rescale=None,
|
||||
show_progress=False, return_intermediate=None, intermediate_callback=None):
|
||||
show_progress=False, return_intermediate=None, intermediate_callback=None, **kwargs):
|
||||
assert isinstance(steps, (int, torch.LongTensor))
|
||||
assert return_intermediate in (None, 'x0', 'xt')
|
||||
assert isinstance(sampler, (str, dict, Config))
|
||||
intermediates = []
|
||||
|
||||
def callback_fn(x_t, t, sigma=None, alpha=None):
|
||||
timestamp = t
|
||||
t = t.repeat(len(x_t)).round().long().to(x_t.device)
|
||||
sigma = sigma.repeat(len(x_t), *([1] * (len(sigma.shape) - 1)))
|
||||
alpha = alpha.repeat(len(x_t), *([1] * (len(alpha.shape) - 1)))
|
||||
|
||||
if guide_scale is None or guide_scale == 1.0:
|
||||
out = model(x=x_t, t=t, **model_kwargs)
|
||||
else:
|
||||
if use_dynamic_cfg:
|
||||
guidance_scale = 1 + guide_scale * ((1 - math.cos(math.pi * ((steps - t.item()) / steps) ** 5.0)) / 2)
|
||||
guidance_scale = 1 + guide_scale * ((1 - math.cos(math.pi * ((steps - timestamp.item()) / steps) ** 5.0)) / 2)
|
||||
else:
|
||||
guidance_scale = guide_scale
|
||||
y_out = model(x=x_t, t=t, **model_kwargs[0])
|
||||
|
||||
@@ -145,12 +145,12 @@ class DDIMSampler(BaseDiffusionSampler):
|
||||
return output
|
||||
|
||||
def step(self, sampler_output):
|
||||
step = sampler_output.step
|
||||
x_t = sampler_output.x_t
|
||||
step = sampler_output.step
|
||||
t = sampler_output.ts[step]
|
||||
sigmas_vp = sampler_output.sigmas_vp.to(x_t.device)
|
||||
alpha_init = _i(sampler_output.alphas_init, step, x_t)
|
||||
sigma_init = _i(sampler_output.sigmas_init, step, x_t)
|
||||
alpha_init = _i(sampler_output.alphas_init, step, x_t[:1])
|
||||
sigma_init = _i(sampler_output.sigmas_init, step, x_t[:1])
|
||||
|
||||
x = sampler_output.callback_fn(x_t, t, sigma_init, alpha_init)
|
||||
noise_factor = self.eta * (sigmas_vp[step + 1] ** 2 / sigmas_vp[step] ** 2 *
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.tuner import sce
|
||||
from scepter.modules.model.tuner.swift_tuner import (SwiftAdapter, SwiftFull,
|
||||
SwiftLoRA)
|
||||
from scepter.modules.model.tuner.swift_tuner import (SwiftPart, SwiftAdapter, SwiftFull,
|
||||
SwiftLoRA, SwiftSCETuning)
|
||||
|
||||
@@ -23,6 +23,29 @@ class SwiftFull(BaseTuner):
|
||||
SwiftFull.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@TUNERS.register_class()
|
||||
class SwiftPart():
|
||||
para_dict = {
|
||||
'TARGET_MODULES': {
|
||||
'value': '',
|
||||
'description': 'The norm expression of target modules.'
|
||||
}
|
||||
}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
from swift.tuners.part import PartConfig
|
||||
self.logger = logger
|
||||
self.init_config = PartConfig(target_modules=cfg.TARGET_MODULES)
|
||||
|
||||
def __call__(self, *args, **kwargs):
|
||||
return self.init_config
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('TUNERS',
|
||||
__class__.__name__,
|
||||
SwiftPart.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@TUNERS.register_class()
|
||||
class SwiftLoRA():
|
||||
|
||||
@@ -293,7 +293,7 @@ class FileSystem(object):
|
||||
wait_finish=wait_finish)
|
||||
else:
|
||||
local_path = None
|
||||
R.acquire(timeout = 2)
|
||||
R.acquire(timeout=60)
|
||||
try:
|
||||
data_quene.put_nowait([target_path, local_path])
|
||||
except Exception:
|
||||
@@ -355,7 +355,7 @@ class FileSystem(object):
|
||||
pass
|
||||
else:
|
||||
flg = False
|
||||
R.acquire(timeout=2)
|
||||
R.acquire(timeout=60)
|
||||
try:
|
||||
data_quene.put_nowait([local_path, target_path, flg])
|
||||
except Exception:
|
||||
|
||||
Reference in New Issue
Block a user