From 3412a074c559f496d8d81d705e4a058552bd6607 Mon Sep 17 00:00:00 2001 From: XDUWQ <1300964705@qq.com> Date: Tue, 25 Jul 2023 15:00:28 +0800 Subject: [PATCH] precommit --- .../custom/finetune_stable_diffusion_custom.py | 3 +-- .../multi_modal/custom_diffusion/custom_diffusion_trainer.py | 5 +++-- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/examples/pytorch/stable_diffusion/custom/finetune_stable_diffusion_custom.py b/examples/pytorch/stable_diffusion/custom/finetune_stable_diffusion_custom.py index 47089bef..007ea82b 100644 --- a/examples/pytorch/stable_diffusion/custom/finetune_stable_diffusion_custom.py +++ b/examples/pytorch/stable_diffusion/custom/finetune_stable_diffusion_custom.py @@ -94,7 +94,6 @@ class StableDiffusionCustomArguments(TrainingArgs): metadata={ 'help': 'Path to json containing multiple concepts.', }) - training_args = StableDiffusionCustomArguments( @@ -160,7 +159,7 @@ pipe = pipeline( task=Tasks.text_to_image_synthesis, model=training_args.model, custom_dir=training_args.work_dir + '/output', - modifier_token='', + modifier_token='+', model_revision=args.model_revision) output = pipe({'text': args.instance_prompt}) diff --git a/modelscope/trainers/multi_modal/custom_diffusion/custom_diffusion_trainer.py b/modelscope/trainers/multi_modal/custom_diffusion/custom_diffusion_trainer.py index 3e6450d7..2421d145 100644 --- a/modelscope/trainers/multi_modal/custom_diffusion/custom_diffusion_trainer.py +++ b/modelscope/trainers/multi_modal/custom_diffusion/custom_diffusion_trainer.py @@ -7,6 +7,7 @@ import warnings from pathlib import Path from typing import Union +import json import numpy as np import torch import torch.nn.functional as F @@ -309,9 +310,9 @@ class CustomDiffusionTrainer(EpochBasedTrainer): 'class_data_dir': self.class_data_dir, }] else: - with open(args.concepts_list, "r") as f: + with open(self.concepts_list, 'r') as f: self.concepts_list = json.load(f) - print("--------self.concepts_list: ", self.concepts_list) + print('--------self.concepts_list: ', self.concepts_list) # Adding a modifier token which is optimized self.modifier_token_id = []