update v1.1.0
This commit is contained in:
@@ -11,6 +11,7 @@ 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.flux_inference import FluxInference
|
||||
from scepter.modules.inference.stylebooth_inference import StyleboothInference
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.distribute import we
|
||||
@@ -226,7 +227,7 @@ class DiffusionInferenceTest(unittest.TestCase):
|
||||
'stylebooth_test_lowpoly_cute_dog.png')
|
||||
save_image(output['images'], save_path)
|
||||
|
||||
# @unittest.skip('')
|
||||
@unittest.skip('')
|
||||
def test_sd3(self):
|
||||
config_file = 'scepter/methods/studio/inference/dit/sd3_pro.yaml'
|
||||
cfg = Config(cfg_file=config_file)
|
||||
@@ -240,6 +241,19 @@ class DiffusionInferenceTest(unittest.TestCase):
|
||||
save_image(output['images'], save_path)
|
||||
print(save_path)
|
||||
|
||||
# @unittest.skip('')
|
||||
def test_flux(self):
|
||||
config_file = 'scepter/methods/studio/inference/dit/flux1.0_dev_pro.yaml'
|
||||
cfg = Config(cfg_file=config_file)
|
||||
diff_infer = FluxInference(logger=self.logger)
|
||||
diff_infer.init_from_cfg(cfg)
|
||||
output = diff_infer({
|
||||
'prompt': '1 girl',
|
||||
'seed': 2024
|
||||
})
|
||||
save_path = os.path.join(self.tmp_dir, 'flux_dev_1girl.png')
|
||||
save_image(output['images'], save_path)
|
||||
print(save_path)
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user