From 8a9b791e3f2396d3b84bbf4d74050962cf374f2b Mon Sep 17 00:00:00 2001 From: LouieStark Date: Tue, 9 Apr 2024 14:02:28 +0800 Subject: [PATCH] fix bug --- readme.md | 1 - .../modules/inference/diffusion_inference.py | 2 +- .../model/network/autoencoder/ae_kl.py | 2 +- .../self_train/self_train_ui/model_ui.py | 32 +++++++++++++------ 4 files changed, 25 insertions(+), 12 deletions(-) diff --git a/readme.md b/readme.md index 685c2d5..2467a43 100644 --- a/readme.md +++ b/readme.md @@ -176,7 +176,6 @@ python run.py --cfg classifier.yaml ``` - ## 🖥️ SCEPTER Studio ### Launch diff --git a/scepter/modules/inference/diffusion_inference.py b/scepter/modules/inference/diffusion_inference.py index 097c045..dd5efa9 100644 --- a/scepter/modules/inference/diffusion_inference.py +++ b/scepter/modules/inference/diffusion_inference.py @@ -395,7 +395,7 @@ class DiffusionInference(): return self.first_stage_model['paras']['scale_factor'] * z def decode_first_stage(self, z): - _, dtype = self.get_function_info(self.first_stage_model, 'encode') + _, dtype = self.get_function_info(self.first_stage_model, 'decode') with torch.autocast('cuda', enabled=dtype == 'float16', dtype=getattr(torch, dtype)): diff --git a/scepter/modules/model/network/autoencoder/ae_kl.py b/scepter/modules/model/network/autoencoder/ae_kl.py index 77bd9e3..8f0c2eb 100644 --- a/scepter/modules/model/network/autoencoder/ae_kl.py +++ b/scepter/modules/model/network/autoencoder/ae_kl.py @@ -52,7 +52,7 @@ class DiagonalGaussianDistribution(object): dim=dims) def mode(self): - print('*** use DiagonalGaussianDistribution.mode() ***') + # print('*** use DiagonalGaussianDistribution.mode() ***') return self.mean diff --git a/scepter/studio/self_train/self_train_ui/model_ui.py b/scepter/studio/self_train/self_train_ui/model_ui.py index d2de96d..47c23e2 100644 --- a/scepter/studio/self_train/self_train_ui/model_ui.py +++ b/scepter/studio/self_train/self_train_ui/model_ui.py @@ -171,16 +171,26 @@ class ModelUI(UIBase): status = trainer_ui.trainer_ins.get_status(model_name) ckpt_list = self.get_ckpt_list(model_name) ckpt_value = ckpt_list[-1] if len(ckpt_list) > 0 else '' + if ckpt_value is not None and len(ckpt_value) > 0: + gallery_value = self.get_gallery_list(model_name, ckpt_value) + else: + gallery_value = [] + select_index = 0 if len(gallery_value) > 0 else None return (message, gr.Column(visible=status in ('running', 'success')), - gr.Dropdown(choices=ckpt_list, value=ckpt_value)) + gr.Dropdown(choices=ckpt_list, value=ckpt_value), + gr.Gallery(value=gallery_value, + preview=True, + selected_index=select_index) + ) self.output_model_name.change(fn=model_name_change, inputs=[self.output_model_name], outputs=[ self.log_message, self.export_log_panel, - self.output_ckpt_name + self.output_ckpt_name, + self.eval_gallery ], queue=False) @@ -243,8 +253,9 @@ class ModelUI(UIBase): ckpt_list = self.get_ckpt_list(model_name) ckpt_value = ckpt_list[-1] if len(ckpt_list) > 0 else '' ret_gallery = ckpt_name_change(model_name, ckpt_value) - return (message, gr.Column(visible=status in ('running', - 'success')), + return (message, + gr.Column(visible=status in ('running', 'success')), + gr.Dropdown(choices=self.model_list, value=model_name), gr.Dropdown(choices=ckpt_list, value=ckpt_value), ret_gallery) @@ -254,7 +265,7 @@ class ModelUI(UIBase): outputs=[ self.log_message, self.export_log_panel, - # self.output_model_name, + self.output_model_name, self.output_ckpt_name, self.eval_gallery ], @@ -305,10 +316,13 @@ class ModelUI(UIBase): image_dir = os.path.join(self.work_dir, output_model, 'eval_probe', output_ckpt_name, 'image') - image_path = [ - os.path.join(image_dir, name) - for name in os.listdir(image_dir) - ] + if os.path.exists(image_dir): + image_path = [ + os.path.join(image_dir, name) + for name in os.listdir(image_dir) + ] + else: + image_path = [] else: _, base_model, base_model_revision, tuner_name, _ = output_model.split( '@', 4)