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 = ''
- 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'
'
- sec_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'