# -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. import os import unittest class TrainTest(unittest.TestCase): def setUp(self): print(('Testing %s.%s' % (type(self).__name__, self._testMethodName))) self.tmp_dir = './cache/save_data' def tearDown(self): super().tearDown() @unittest.skip('') def test_generation_example_full(self): os.system( 'python scepter/tools/run_train.py ' '--cfg scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml ' '--max_steps 100') self.assertTrue( os.path.exists( os.path.join(self.tmp_dir, 'sd15_512_full/checkpoints'))) os.system( 'python scepter/tools/run_train.py ' '--cfg scepter/methods/examples/generation/stable_diffusion_2.1_768.yaml ' '--max_steps 100') self.assertTrue( os.path.exists( os.path.join(self.tmp_dir, 'sd21_768_full/checkpoints'))) os.system( 'python scepter/tools/run_train.py ' '--cfg scepter/methods/examples/generation/stable_diffusion_xl_1024.yaml ' '--max_steps 100') self.assertTrue( os.path.exists( os.path.join(self.tmp_dir, 'sdxl_1024_full/checkpoints'))) @unittest.skip('') def test_generation_example_lora(self): os.system( 'python scepter/tools/run_train.py ' '--cfg scepter/methods/examples/generation/stable_diffusion_1.5_512_lora.yaml ' '--max_steps 100') self.assertTrue( os.path.exists( os.path.join(self.tmp_dir, 'sd15_512_lora/checkpoints'))) os.system( 'python scepter/tools/run_train.py ' '--cfg scepter/methods/examples/generation/stable_diffusion_2.1_768_lora.yaml ' '--max_steps 100') self.assertTrue( os.path.exists( os.path.join(self.tmp_dir, 'sd21_768_lora/checkpoints'))) os.system( 'python scepter/tools/run_train.py ' '--cfg scepter/methods/examples/generation/stable_diffusion_xl_1024_lora.yaml ' '--max_steps 100') self.assertTrue( os.path.exists( os.path.join(self.tmp_dir, 'sdxl_1024_lora/checkpoints'))) @unittest.skip('') def test_generation_example_scedit_t2i_swift(self): os.system( 'python scepter/tools/run_train.py ' '--cfg scepter/methods/scedit/t2i/sd15_512_sce_t2i_swift.yaml ' '--max_steps 100') self.assertTrue( os.path.exists( os.path.join(self.tmp_dir, 'sd15_512_sce_t2i_swift/checkpoints'))) os.system( 'python scepter/tools/run_train.py ' '--cfg scepter/methods/scedit/t2i/sd21_768_sce_t2i_swift.yaml ' '--max_steps 100') self.assertTrue( os.path.exists( os.path.join(self.tmp_dir, 'sd21_768_sce_t2i_swift/checkpoints'))) os.system( 'python scepter/tools/run_train.py ' '--cfg scepter/methods/scedit/t2i/sdxl_1024_sce_t2i_swift.yaml ' '--max_steps 100') self.assertTrue( os.path.exists( os.path.join(self.tmp_dir, 'sdxl_1024_sce_t2i_swift/checkpoints'))) @unittest.skip('') def test_generation_example_scedit_t2i(self): os.system('python scepter/tools/run_train.py ' '--cfg scepter/methods/scedit/t2i/sd15_512_sce_t2i.yaml ' '--max_steps 100') self.assertTrue( os.path.exists( os.path.join(self.tmp_dir, 'sd15_512_sce_t2i/checkpoints'))) os.system('python scepter/tools/run_train.py ' '--cfg scepter/methods/scedit/t2i/sd21_768_sce_t2i.yaml ' '--max_steps 100') self.assertTrue( os.path.exists( os.path.join(self.tmp_dir, 'sd21_768_sce_t2i/checkpoints'))) os.system('python scepter/tools/run_train.py ' '--cfg scepter/methods/scedit/t2i/sdxl_1024_sce_t2i.yaml ' '--max_steps 100') self.assertTrue( os.path.exists( os.path.join(self.tmp_dir, 'sdxl_1024_sce_t2i/checkpoints'))) @unittest.skip('') def test_generation_example_scedit_ctr(self): os.system('python scepter/tools/run_train.py ' '--cfg scepter/methods/scedit/ctr/sd15_512_sce_ctr_hed.yaml ' '--max_steps 100') self.assertTrue( os.path.exists( os.path.join(self.tmp_dir, 'sd15_512_sce_ctr_hed/checkpoints'))) os.system( 'python scepter/tools/run_train.py ' '--cfg scepter/methods/scedit/ctr/sd21_768_sce_ctr_canny.yaml ' '--max_steps 100') self.assertTrue( os.path.exists( os.path.join(self.tmp_dir, 'sd21_768_sce_ctr_canny/checkpoints'))) os.system( 'python scepter/tools/run_train.py ' '--cfg scepter/methods/scedit/ctr/sd21_768_sce_ctr_pose.yaml ' '--max_steps 100') self.assertTrue( os.path.exists( os.path.join(self.tmp_dir, 'sd21_768_sce_ctr_pose/checkpoints'))) os.system( 'python scepter/tools/run_train.py ' '--cfg scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_depth.yaml ' '--max_steps 100') self.assertTrue( os.path.exists( os.path.join(self.tmp_dir, 'sdxl_1024_sce_ctr_depth/checkpoints'))) os.system( 'python scepter/tools/run_train.py ' '--cfg scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_color.yaml ' '--max_steps 100') self.assertTrue( os.path.exists( os.path.join(self.tmp_dir, 'sdxl_1024_sce_ctr_color/checkpoints'))) # @unittest.skip('') def test_generation_example_datatxt(self): # os.system( # 'python scepter/tools/run_train.py ' # '--cfg scepter/methods/scedit/t2i/sdxl_1024_sce_t2i_datatxt.yaml ' # '--max_steps 100' # ) # self.assertTrue(os.path.exists(os.path.join(self.tmp_dir, 'sdxl_1024_sce_t2i_datatxt/checkpoints'))) os.system( 'python scepter/tools/run_train.py ' '--cfg scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_color_datatxt.yaml ' '--max_steps 100') self.assertTrue( os.path.exists( os.path.join(self.tmp_dir, 'sdxl_1024_sce_ctr_color_datatxt/checkpoints'))) if __name__ == '__main__': unittest.main()