v1.0.3 update

This commit is contained in:
zeyinzi.jzyz
2024-07-18 14:12:42 +08:00
parent 7a9f90efb2
commit 01fd8335af
94 changed files with 8776 additions and 417 deletions
+16 -1
View File
@@ -10,6 +10,7 @@ import torchvision.transforms.functional as TF
from PIL import Image
from scepter.modules.annotator.registry import ANNOTATORS
from scepter.modules.inference.diffusion_inference import DiffusionInference
from scepter.modules.inference.sd3_inference import SD3Inference
from scepter.modules.inference.stylebooth_inference import StyleboothInference
from scepter.modules.utils.config import Config
from scepter.modules.utils.distribute import we
@@ -98,7 +99,7 @@ class DiffusionInferenceTest(unittest.TestCase):
save_path = os.path.join(self.tmp_dir, 'sd15_flower_2d.png')
save_image(output['images'], save_path)
# @unittest.skip('')
@unittest.skip('')
def test_sd21_scedit_ctr_canny(self):
# init model
config_file = 'scepter/methods/studio/inference/stable_diffusion/sd21_pro.yaml'
@@ -225,6 +226,20 @@ class DiffusionInferenceTest(unittest.TestCase):
'stylebooth_test_lowpoly_cute_dog.png')
save_image(output['images'], save_path)
# @unittest.skip('')
def test_sd3(self):
config_file = 'scepter/methods/studio/inference/dit/sd3_pro.yaml'
cfg = Config(cfg_file=config_file)
diff_infer = SD3Inference(logger=self.logger)
diff_infer.init_from_cfg(cfg)
output = diff_infer({
'prompt': 'a cat holds a blackboard that writes "hello world"',
'seed': 2024
})
save_path = os.path.join(self.tmp_dir, 'sd3_cat.png')
save_image(output['images'], save_path)
print(save_path)
if __name__ == '__main__':
unittest.main()