diff --git a/scepter/methods/studio/chatbot/models/ace_0.6b_1024.yaml b/scepter/methods/studio/chatbot/models/ace_0.6b_1024.yaml index 668dac2..c988e5f 100644 --- a/scepter/methods/studio/chatbot/models/ace_0.6b_1024.yaml +++ b/scepter/methods/studio/chatbot/models/ace_0.6b_1024.yaml @@ -1,6 +1,6 @@ NAME: ACE_0.6B_1024 IS_DEFAULT: False -USE_DYNAMIC_MODEL: False +USE_DYNAMIC_MODEL: True DEFAULT_PARAS: PARAS: # diff --git a/scepter/methods/studio/chatbot/models/ace_0.6b_1024_refiner.yaml b/scepter/methods/studio/chatbot/models/ace_0.6b_1024_refiner.yaml index 5647bc1..94f04a4 100644 --- a/scepter/methods/studio/chatbot/models/ace_0.6b_1024_refiner.yaml +++ b/scepter/methods/studio/chatbot/models/ace_0.6b_1024_refiner.yaml @@ -1,6 +1,6 @@ NAME: ACE_0.6B_1024_REFINER IS_DEFAULT: False -USE_DYNAMIC_MODEL: False +USE_DYNAMIC_MODEL: True DEFAULT_PARAS: PARAS: # diff --git a/scepter/methods/studio/chatbot/models/ace_0.6b_512.yaml b/scepter/methods/studio/chatbot/models/ace_0.6b_512.yaml index 0872224..c803d67 100644 --- a/scepter/methods/studio/chatbot/models/ace_0.6b_512.yaml +++ b/scepter/methods/studio/chatbot/models/ace_0.6b_512.yaml @@ -1,6 +1,6 @@ NAME: ACE_0.6B_512 IS_DEFAULT: True -USE_DYNAMIC_MODEL: False +USE_DYNAMIC_MODEL: True DEFAULT_PARAS: PARAS: # diff --git a/scepter/modules/inference/ace_inference.py b/scepter/modules/inference/ace_inference.py index fa8b37e..fd1769d 100644 --- a/scepter/modules/inference/ace_inference.py +++ b/scepter/modules/inference/ace_inference.py @@ -156,17 +156,17 @@ class RefinerInference(DiffusionInference): noise.append(noise_) noise, x_shapes = pack_imagelist_into_tensor(noise) if reverse_scale > 0: - if self.use_dynamic_model: self.dynamic_load(self.first_stage_model, 'first_stage_model') + self.dynamic_load(self.first_stage_model, 'first_stage_model') x_samples = [x.unsqueeze(0) for x in x_samples] x_start = self.encode_first_stage(x_samples, **kwargs) - if self.use_dynamic_model: self.dynamic_unload(self.first_stage_model, + self.dynamic_unload(self.first_stage_model, 'first_stage_model', - skip_loaded=True) + skip_loaded=not self.use_dynamic_model) x_start, _ = pack_imagelist_into_tensor(x_start) else: x_start = None # cond stage - if self.use_dynamic_model: self.dynamic_load(self.cond_stage_model, 'cond_stage_model') + self.dynamic_load(self.cond_stage_model, 'cond_stage_model') function_name, dtype = self.get_function_info(self.cond_stage_model) with torch.autocast('cuda', enabled=dtype == 'float16', @@ -174,12 +174,12 @@ class RefinerInference(DiffusionInference): ctx = getattr(get_model(self.cond_stage_model), function_name)(prompt) ctx["x_shapes"] = x_shapes - if self.use_dynamic_model: self.dynamic_unload(self.cond_stage_model, + self.dynamic_unload(self.cond_stage_model, 'cond_stage_model', - skip_loaded=True) + skip_loaded=not self.use_dynamic_model) - if self.use_dynamic_model: self.dynamic_load(self.diffusion_model, 'diffusion_model') + self.dynamic_load(self.diffusion_model, 'diffusion_model') # UNet use input n_prompt function_name, dtype = self.get_function_info( self.diffusion_model) @@ -207,14 +207,14 @@ class RefinerInference(DiffusionInference): x=x_start, **kwargs).float() latent = unpack_tensor_into_imagelist(latent, x_shapes) - if self.use_dynamic_model: self.dynamic_unload(self.diffusion_model, + self.dynamic_unload(self.diffusion_model, 'diffusion_model', - skip_loaded=True) - if self.use_dynamic_model: self.dynamic_load(self.first_stage_model, 'first_stage_model') + skip_loaded=not self.use_dynamic_model) + self.dynamic_load(self.first_stage_model, 'first_stage_model') x_samples = self.decode_first_stage(latent) - if self.use_dynamic_model: self.dynamic_unload(self.first_stage_model, + self.dynamic_unload(self.first_stage_model, 'first_stage_model', - skip_loaded=True) + skip_loaded=not self.use_dynamic_model) return x_samples @@ -398,11 +398,11 @@ class ACEInference(DiffusionInference): if use_ace and (not is_txt_image or refiner_scale <= 0): ctx, null_ctx = {}, {} # Get Noise Shape - if self.use_dynamic_model: self.dynamic_load(self.first_stage_model, 'first_stage_model') + self.dynamic_load(self.first_stage_model, 'first_stage_model') x = self.encode_first_stage(image) - if self.use_dynamic_model: self.dynamic_unload(self.first_stage_model, + self.dynamic_unload(self.first_stage_model, 'first_stage_model', - skip_loaded=True) + skip_loaded=not self.use_dynamic_model) noise = [ torch.empty(*i.shape, device=we.device_id).normal_(generator=g) for i in x @@ -416,7 +416,7 @@ class ACEInference(DiffusionInference): ctx['x_mask'] = null_ctx['x_mask'] = cond_mask # Encode Prompt - if self.use_dynamic_model: self.dynamic_load(self.cond_stage_model, 'cond_stage_model') + self.dynamic_load(self.cond_stage_model, 'cond_stage_model') function_name, dtype = self.get_function_info(self.cond_stage_model) cont, cont_mask = getattr(get_model(self.cond_stage_model), function_name)(prompt) @@ -426,14 +426,14 @@ class ACEInference(DiffusionInference): function_name)(n_prompt) null_cont, null_cont_mask = self.cond_stage_embeddings( prompt, edit_image, null_cont, null_cont_mask) - if self.use_dynamic_model: self.dynamic_unload(self.cond_stage_model, + self.dynamic_unload(self.cond_stage_model, 'cond_stage_model', - skip_loaded=False) + skip_loaded=not self.use_dynamic_model) ctx['crossattn'] = cont null_ctx['crossattn'] = null_cont # Encode Edit Images - if self.use_dynamic_model: self.dynamic_load(self.first_stage_model, 'first_stage_model') + self.dynamic_load(self.first_stage_model, 'first_stage_model') edit_image = [to_device(i, strict=False) for i in edit_image] edit_image_mask = [to_device(i, strict=False) for i in edit_image_mask] e_img, e_mask = [], [] @@ -444,14 +444,14 @@ class ACEInference(DiffusionInference): m = [None] * len(u) e_img.append(self.encode_first_stage(u, **kwargs)) e_mask.append([self.interpolate_func(i) for i in m]) - if self.use_dynamic_model: self.dynamic_unload(self.first_stage_model, + self.dynamic_unload(self.first_stage_model, 'first_stage_model', - skip_loaded=True) + skip_loaded=not self.use_dynamic_model) null_ctx['edit'] = ctx['edit'] = e_img null_ctx['edit_mask'] = ctx['edit_mask'] = e_mask # Diffusion Process - if self.use_dynamic_model: self.dynamic_load(self.diffusion_model, 'diffusion_model') + self.dynamic_load(self.diffusion_model, 'diffusion_model') function_name, dtype = self.get_function_info(self.diffusion_model) with torch.autocast('cuda', enabled=dtype in ('float16', 'bfloat16'), @@ -492,17 +492,17 @@ class ACEInference(DiffusionInference): guide_rescale=guide_rescale, return_intermediate=None, **kwargs) - if self.use_dynamic_model: self.dynamic_unload(self.diffusion_model, + self.dynamic_unload(self.diffusion_model, 'diffusion_model', - skip_loaded=False) + skip_loaded=not self.use_dynamic_model) # Decode to Pixel Space - if self.use_dynamic_model: self.dynamic_load(self.first_stage_model, 'first_stage_model') + self.dynamic_load(self.first_stage_model, 'first_stage_model') samples = unpack_tensor_into_imagelist(latent, x_shapes) x_samples = self.decode_first_stage(samples) - if self.use_dynamic_model: self.dynamic_unload(self.first_stage_model, + self.dynamic_unload(self.first_stage_model, 'first_stage_model', - skip_loaded=False) + skip_loaded=not self.use_dynamic_model) x_samples = [x.squeeze(0) for x in x_samples] else: x_samples = image diff --git a/scepter/modules/solver/diffusion_solver.py b/scepter/modules/solver/diffusion_solver.py index 3f5c14a..2d90d7b 100644 --- a/scepter/modules/solver/diffusion_solver.py +++ b/scepter/modules/solver/diffusion_solver.py @@ -197,7 +197,7 @@ class LatentDiffusionSolver(BaseSolver): else: self.logger.info('Use default backend.') self.use_scaler = cfg.get('USE_SCALER', True) - self.enable_gradscaler = cfg.get('ENABLE_GRADSCALER', True) + self.enable_gradscaler = cfg.get('ENABLE_GRADSCALER', False) self.use_orig_params = cfg.get('USE_ORIG_PARAMS', False) self.model_shard = cfg.get('SHARDING_STRATEGY', 'full_shard') self.sharding_size = cfg.get('SHARDING_SIZE', None) @@ -422,8 +422,10 @@ class LatentDiffusionSolver(BaseSolver): process_group=None) else: self.scaler = amp.GradScaler(enabled=self.enable_gradscaler) - else: + elif self.cfg.DTYPE in ['float16']: self.scaler = amp.GradScaler() + else: + self.scaler = None else: self.scaler = None diff --git a/scepter/modules/utils/visualization.py b/scepter/modules/utils/visualization.py index 8f2810d..82ba90f 100644 --- a/scepter/modules/utils/visualization.py +++ b/scepter/modules/utils/visualization.py @@ -1,6 +1,10 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +import copy from enum import Enum +import os + +from scepter.modules.utils.file_system import FS class Media(Enum): @@ -13,15 +17,17 @@ class Media(Enum): class HtmlVisualization(object): - def __init__(self, - allow_annotation=False, - slice_size=1000, - align='center', - width_scale='60%', - title='Visualization', - height=600, - width=None, - text_cols=40): + def __init__( + self, + allow_annotation=False, + slice_size=1000, + align='center', + width_scale='60%', + title='Visualization', + height=600, + width=None, + text_cols=40 + ): self.content_list = [] self.rows_meta = [] self.allow_annotation = allow_annotation @@ -31,9 +37,9 @@ class HtmlVisualization(object): self.title = title self.html_start = '' self.html_head = f'{title}' - self.height = height if height is not None else '600' - self.width = width if width is not None else 'auto' - self.text_cols = text_cols if text_cols is not None else 'auto' + self.height = height if height is not None else "600" + self.width = width if width is not None else "auto" + self.text_cols = text_cols if text_cols is not None else "auto" self.html_style = (''' \n \n - '''.replace('{width_scale}', self.width_scale).replace( - '{align}', self.align).replace('{pair_height}', f'{self.height}')) + '''.replace('{width_scale}', + self.width_scale).replace('{align}', self.align) + .replace('{pair_height}', f'{self.height}')) self.html_body_script = ''' \n @@ -152,16 +165,17 @@ class HtmlVisualization(object): ''' self.label_button = ( - '
' + - "" - + '
') + '
' + + "" + + '
') def format_col(self, content='', label='', type=Media.TEXT, show_label=True, - cols_span=1): + cols_span=1 + ): if type == Media.TEXT: ret_str = '"{content}"' - sec_ret_str = f'{label}' if show_label else '' + sec_ret_str = f'{label}' if show_label else "" elif type == Media.IMAGE: ret_str = f'{label}' if show_label else '' + sec_ret_str = f'{label}' if show_label else "" elif type == Media.VIDEO: ret_str = '' - sec_ret_str = f'{label}' if show_label else '' + sec_ret_str = f'{label}' if show_label else "" elif type == Media.AUDIO: ret_str = f'