mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
+10
-5
@@ -7,7 +7,6 @@ from modules import scripts, processing, shared, sd_models
|
||||
class Script(scripts.Script):
|
||||
supported_models = ['sd', 'sdxl']
|
||||
orig_pipe = None
|
||||
library = None
|
||||
|
||||
def title(self):
|
||||
return 'K-Diffusion'
|
||||
@@ -34,6 +33,8 @@ class Script(scripts.Script):
|
||||
_step = d['i']
|
||||
|
||||
def run(self, p: processing.StableDiffusionProcessing, sampler: str): # pylint: disable=arguments-differ
|
||||
if sampler is None or len(sampler) == 0:
|
||||
return None
|
||||
if shared.sd_model_type not in self.supported_models:
|
||||
shared.log.warning(f'K-Diffusion: class={shared.sd_model.__class__.__name__} model={shared.sd_model_type} required={self.supported_models}')
|
||||
return None
|
||||
@@ -44,12 +45,16 @@ class Script(scripts.Script):
|
||||
cls = diffusers.pipelines.StableDiffusionXLKDiffusionPipeline
|
||||
if cls is None:
|
||||
return None
|
||||
from modules import sd_samplers_kdiffusion
|
||||
|
||||
sampler_fn = getattr(sd_samplers_kdiffusion.k_sampling, f'sample_{sampler}', None)
|
||||
if sampler_fn is None:
|
||||
shared.log.warning(f'K-Diffusion: sampler={sampler} not found')
|
||||
return None
|
||||
|
||||
self.orig_pipe = shared.sd_model
|
||||
shared.sd_model = sd_models.switch_pipe(cls, shared.sd_model)
|
||||
sampler = 'sample_' + sampler
|
||||
|
||||
sampling = getattr(self.library, "sampling", None)
|
||||
shared.sd_model.sampler = getattr(sampling, sampler)
|
||||
shared.sd_model.sampler = sampler_fn
|
||||
|
||||
params = inspect.signature(shared.sd_model.sampler).parameters.values()
|
||||
params = {param.name: param.default for param in params if param.default != inspect.Parameter.empty}
|
||||
|
||||
Reference in New Issue
Block a user