v1.0.0 update
This commit is contained in:
@@ -8,8 +8,6 @@ import numpy as np
|
||||
import torchvision.transforms as TT
|
||||
import torchvision.transforms.functional as TF
|
||||
from PIL import Image
|
||||
from torchvision.utils import save_image
|
||||
|
||||
from scepter.modules.annotator.registry import ANNOTATORS
|
||||
from scepter.modules.inference.diffusion_inference import DiffusionInference
|
||||
from scepter.modules.inference.stylebooth_inference import StyleboothInference
|
||||
@@ -17,6 +15,7 @@ from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.modules.utils.logger import get_logger
|
||||
from torchvision.utils import save_image
|
||||
|
||||
|
||||
class DiffusionInferenceTest(unittest.TestCase):
|
||||
@@ -68,7 +67,7 @@ class DiffusionInferenceTest(unittest.TestCase):
|
||||
'sdxl_test_prompt_a_cute_dog.png')
|
||||
save_image(output['images'], save_path)
|
||||
|
||||
# @unittest.skip('')
|
||||
@unittest.skip('')
|
||||
def test_sd15_scedit_t2i_2D(self):
|
||||
# init model
|
||||
config_file = 'scepter/methods/studio/inference/stable_diffusion/sd15_pro.yaml'
|
||||
@@ -84,7 +83,7 @@ class DiffusionInferenceTest(unittest.TestCase):
|
||||
'IMAGE_PATH': None,
|
||||
'TUNER_TYPE': 'SwiftSCE',
|
||||
'MODEL_PATH':
|
||||
'ms://damo/scepter_scedit@tuners_model/SD1.5/Flat2DArt',
|
||||
'ms://iic/scepter_scedit@tuners_model/SD1.5/Flat2DArt',
|
||||
'PROMPT_EXAMPLE': None
|
||||
}
|
||||
tuner_model = Config(cfg_dict=tuner_model, load=False)
|
||||
@@ -100,6 +99,61 @@ class DiffusionInferenceTest(unittest.TestCase):
|
||||
save_image(output['images'], save_path)
|
||||
|
||||
# @unittest.skip('')
|
||||
def test_sd21_scedit_ctr_canny(self):
|
||||
# init model
|
||||
config_file = 'scepter/methods/studio/inference/stable_diffusion/sd21_pro.yaml'
|
||||
cfg = Config(cfg_file=config_file)
|
||||
diff_infer = DiffusionInference(logger=self.logger)
|
||||
diff_infer.init_from_cfg(cfg)
|
||||
# extract condition
|
||||
canny_dict = {
|
||||
'NAME': 'CannyAnnotator',
|
||||
'LOW_THRESHOLD': 100,
|
||||
'HIGH_THRESHOLD': 200
|
||||
}
|
||||
canny_anno = Config(cfg_dict=canny_dict, load=False)
|
||||
canny_ins = ANNOTATORS.build(canny_anno).to(we.device_id)
|
||||
output_height, output_width = 768, 768
|
||||
control_cond_image = Image.open('asset/images/flower.jpg')
|
||||
control_cond_image = TT.Resize(min(output_height,
|
||||
output_width))(control_cond_image)
|
||||
control_cond_image = TT.CenterCrop(
|
||||
(output_height, output_width))(control_cond_image)
|
||||
control_cond_image = canny_ins(np.array(control_cond_image))
|
||||
control_save_path = os.path.join(self.tmp_dir,
|
||||
'sd21_flower_canny_preproccess.png')
|
||||
save_image(TF.to_tensor(control_cond_image), control_save_path)
|
||||
control_cond_image = Image.open(control_save_path)
|
||||
# load control model
|
||||
control_model = {
|
||||
'NAME':
|
||||
'canny',
|
||||
'NAME_ZH':
|
||||
None,
|
||||
'DESCRIPTION':
|
||||
None,
|
||||
'BASE_MODEL':
|
||||
'SD2.1',
|
||||
'TYPE':
|
||||
'Canny',
|
||||
'MODEL_PATH':
|
||||
'ms://iic/scepter_scedit@controllable_model/SD2.1/canny_control'
|
||||
}
|
||||
control_model = Config(cfg_dict=control_model, load=False)
|
||||
# prepare data
|
||||
input_data = {'prompt': 'a single flower is shown in front of a tree'}
|
||||
input_params = {
|
||||
'control_model': control_model,
|
||||
'control_cond_image': control_cond_image,
|
||||
'control_scale': 1.0,
|
||||
'crop_type': 'CenterCrop',
|
||||
'seed': 2024
|
||||
}
|
||||
output = diff_infer(input_data, **input_params)
|
||||
save_path = os.path.join(self.tmp_dir, 'sd21_flower_canny.png')
|
||||
save_image(output['images'], save_path)
|
||||
|
||||
@unittest.skip('')
|
||||
def test_sdxl_scedit_ctr_canny(self):
|
||||
# init model
|
||||
config_file = 'scepter/methods/studio/inference/sdxl/sdxl1.0_pro.yaml'
|
||||
@@ -138,7 +192,7 @@ class DiffusionInferenceTest(unittest.TestCase):
|
||||
'TYPE':
|
||||
'Canny',
|
||||
'MODEL_PATH':
|
||||
'ms://damo/scepter_scedit@controllable_model/SD_XL1.0/canny_control'
|
||||
'ms://iic/scepter_scedit@controllable_model/SD_XL1.0/canny_control'
|
||||
}
|
||||
control_model = Config(cfg_dict=control_model, load=False)
|
||||
# prepare data
|
||||
@@ -154,7 +208,7 @@ class DiffusionInferenceTest(unittest.TestCase):
|
||||
save_path = os.path.join(self.tmp_dir, 'sdxl_flower_canny.png')
|
||||
save_image(output['images'], save_path)
|
||||
|
||||
# @unittest.skip('')
|
||||
@unittest.skip('')
|
||||
def test_stylebooth(self):
|
||||
config_file = 'scepter/methods/studio/inference/edit/stylebooth_tb_pro.yaml'
|
||||
cfg = Config(cfg_file=config_file)
|
||||
@@ -166,7 +220,7 @@ class DiffusionInferenceTest(unittest.TestCase):
|
||||
style_edit_image=Image.open(
|
||||
'asset/images/inpainting_text_ref/ex4_scene_im.jpg'),
|
||||
style_guide_scale_text=7.5,
|
||||
style_guide_scale_image=0.5)
|
||||
style_guide_scale_image=1.5)
|
||||
save_path = os.path.join(self.tmp_dir,
|
||||
'stylebooth_test_lowpoly_cute_dog.png')
|
||||
save_image(output['images'], save_path)
|
||||
|
||||
Reference in New Issue
Block a user