From 568bd9dac72eed8bde67268ca69ee76f4901dfb3 Mon Sep 17 00:00:00 2001 From: XDUWQ <1300964705@qq.com> Date: Wed, 12 Jul 2023 10:09:37 +0800 Subject: [PATCH] custom diffusion --- .../stable_diffusion_pipeline.py | 25 ++++++++++++++++++- 1 file changed, 24 insertions(+), 1 deletion(-) diff --git a/modelscope/pipelines/multi_modal/diffusers_wrapped/stable_diffusion/stable_diffusion_pipeline.py b/modelscope/pipelines/multi_modal/diffusers_wrapped/stable_diffusion/stable_diffusion_pipeline.py index f09d459d..61a233fa 100644 --- a/modelscope/pipelines/multi_modal/diffusers_wrapped/stable_diffusion/stable_diffusion_pipeline.py +++ b/modelscope/pipelines/multi_modal/diffusers_wrapped/stable_diffusion/stable_diffusion_pipeline.py @@ -25,12 +25,27 @@ from modelscope.utils.constant import Tasks module_name=Pipelines.diffusers_stable_diffusion) class StableDiffusionPipeline(DiffusersPipeline): - def __init__(self, model: str, lora_dir: str = None, **kwargs): + def __init__(self, + model: str, + lora_dir: str = None, + custom_dir: str = None, + modifier_token: str = None, + **kwargs): """ use `model` to create a stable diffusion pipeline Args: model: model id on modelscope hub or local model dir. + lora_dir: lora weight dir for unet. + custom_dir: custom diffusion weight dir for unet. + modifier_token: token to use as a modifier for the concept of custom diffusion. """ + # check custom diffusion input value + if custom_dir is None and modifier_token is not None: + raise ValueError( + 'custom_dir is None but modifier_token is not None') + elif custom_dir is not None and modifier_token is None: + raise ValueError( + 'modifier_token is None but custom_dir is not None') self.device = 'cuda' if torch.cuda.is_available() else 'cpu' # load pipeline @@ -38,10 +53,18 @@ class StableDiffusionPipeline(DiffusersPipeline): self.pipeline = DiffuserStableDiffusionPipeline.from_pretrained( model, torch_dtype=torch_type) self.pipeline = self.pipeline.to(self.device) + # load lora moudle to unet if lora_dir is not None: assert os.path.exists(lora_dir), f"{lora_dir} isn't exist" self.pipeline.unet.load_attn_procs(lora_dir) + # load custom diffusion to unet + if custom_dir is not None: + assert os.path.exists(custom_dir), f"{custom_dir} isn't exist" + self.pipeline.unet.load_attn_procs( + custom_dir, weight_name='pytorch_custom_diffusion_weights.bin') + self.pipeline.load_textual_inversion( + custom_dir, weight_name=f'{modifier_token}.bin') def preprocess(self, inputs: Dict[str, Any], **kwargs) -> Dict[str, Any]: return inputs