This commit is contained in:
jiangzeyinzi
2024-10-23 10:15:31 +08:00
parent 342c6b8a15
commit 379b94ab4f
6 changed files with 58 additions and 9 deletions
@@ -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])
+3 -3
View File
@@ -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 *
+2 -2
View File
@@ -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():
+2 -2
View File
@@ -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: