diff --git a/readme.md b/readme.md
index ffdb186..4e3e3cd 100644
--- a/readme.md
+++ b/readme.md
@@ -179,6 +179,42 @@ PYTHONPATH=. python scepter/tools/webui.py --cfg scepter/methods/studio/scepter_
We deploy a work studio on Modelscope that includes only the inference tab, please refer to [ms_scepter_studio](https://www.modelscope.cn/studios/damo/scepter_studio/summary)
+## 🖼️ Gallery
+
+### Dragon Year Special: Dragon Tuner
+
+
+
+ | Gold Dragon Tuner |
+ Sloppy Dragon Tuner |
+ Red Dragon Tuner + Papercraft Mantra |
+ Azure Dragon Tuner + Pose Control |
+
+
+  |
+  |
+  |
+  |
+
+
+
+### Text Effect Image
+
+
+
+ | Conditional Image |
+ Midas Control "Race track, top view" |
+ Midas Control + Watercolor Mantra "white lilies" |
+ Midas Control + Dragon Tuner "Spring Festival, Chinese dragon" |
+
+
+  |
+  |
+  |
+  |
+
+
+
## ✨ Features
### Text-to-Image Generation
@@ -204,9 +240,9 @@ We deploy a work studio on Modelscope that includes only the inference tab, plea
- 🪄 denotes that the model has been published.
- More models will be released in the future.
-| Model | URL |
-|--------|------------------------------------------------------------------------------------------------------------------------------------------------|
-| SCEdit | [ModelScope](https://modelscope.cn/models/damo/scepter_scedit/summary) [HuggingFace](https://huggingface.co/scepter-studio/scepter_scedit) |
+| Model | URL |
+|--------|-------------------------------------------------------------------------------------|
+| SCEdit | [ModelCard](https://modelscope.cn/models/damo/scepter_scedit/summary) |
PS: Scripts running within the SCEPTER framework will automatically fetch and load models based on the required dependency files, eliminating the need for manual downloads.
diff --git a/scepter/modules/model/backbone/unet/unet_module.py b/scepter/modules/model/backbone/unet/unet_module.py
index 7900196..6aef1fe 100644
--- a/scepter/modules/model/backbone/unet/unet_module.py
+++ b/scepter/modules/model/backbone/unet/unet_module.py
@@ -492,13 +492,20 @@ class DiffusionUNet(BaseModel):
hs.append(h)
h = self.middle_block(h, emb, context)
for m_id, module in enumerate(self.output_blocks):
- h = torch.cat([h, self.lsc_identity[m_id](hs.pop())], dim=1)
+ skip_h = hs.pop()
+ if 'tuner_scale' in kwargs and kwargs[
+ 'tuner_scale'] is not None and kwargs['tuner_scale'] < 1.0:
+ tuner_scale = kwargs['tuner_scale']
+ tuner_h = self.lsc_identity[m_id](skip_h) - skip_h
+ h = torch.cat([h, skip_h + tuner_scale * tuner_h], dim=1)
+ else:
+ h = torch.cat([h, self.lsc_identity[m_id](skip_h)], dim=1)
target_size = hs[-1].shape[-2:] if len(hs) > 0 else None
h = module(h, emb, context, target_size)
out = self.out(h)
return out
- def _forward_control(self, x, emb, context, hint, alpha=0.5, **kwargs):
+ def _forward_control(self, x, emb, context, hint, **kwargs):
control_scale = kwargs.pop('control_scale', 1.0)
multi_csc_tuners = self.control_blocks
# hints
@@ -535,8 +542,8 @@ class DiffusionUNet(BaseModel):
skip_h_new = skip_h + control_scale * multi_control_h
else:
# csc-tuner + sc-tuner
- skip_h_new = skip_h + alpha * multi_control_h + (
- 1 - alpha) * tuner_h
+ tuner_scale = kwargs['tuner_scale']
+ skip_h_new = skip_h + control_scale * multi_control_h + tuner_scale * tuner_h
h = torch.cat([h, skip_h_new], dim=1)
target_size = hs[-1].shape[-2:] if len(hs) > 0 else None
h = module(h, emb, context, target_size)
@@ -834,13 +841,20 @@ class DiffusionUNetXL(DiffusionUNet):
hs.append(h)
h = self.middle_block(h, emb, context)
for m_id, module in enumerate(self.output_blocks):
- h = torch.cat([h, self.lsc_identity[m_id](hs.pop())], dim=1)
+ skip_h = hs.pop()
+ if 'tuner_scale' in kwargs and kwargs[
+ 'tuner_scale'] is not None and kwargs['tuner_scale'] < 1.0:
+ tuner_scale = kwargs['tuner_scale']
+ tuner_h = self.lsc_identity[m_id](skip_h) - skip_h
+ h = torch.cat([h, skip_h + tuner_scale * tuner_h], dim=1)
+ else:
+ h = torch.cat([h, self.lsc_identity[m_id](skip_h)], dim=1)
target_size = hs[-1].shape[-2:] if len(hs) > 0 else None
h = module(h, emb, context, target_size)
out = self.out(h)
return out
- def _forward_control(self, x, emb, context, hint, alpha=0.5, **kwargs):
+ def _forward_control(self, x, emb, context, hint, **kwargs):
control_scale = kwargs.pop('control_scale', 1.0)
multi_csc_tuners = self.control_blocks
# hints
@@ -877,8 +891,8 @@ class DiffusionUNetXL(DiffusionUNet):
skip_h_new = skip_h + control_scale * multi_control_h
else:
# csc-tuner + sc-tuner
- skip_h_new = skip_h + alpha * multi_control_h + (
- 1 - alpha) * tuner_h
+ tuner_scale = kwargs['tuner_scale']
+ skip_h_new = skip_h + control_scale * multi_control_h + tuner_scale * tuner_h
h = torch.cat([h, skip_h_new], dim=1)
target_size = hs[-1].shape[-2:] if len(hs) > 0 else None
h = module(h, emb, context, target_size)
diff --git a/scepter/modules/model/network/diffusion/diffusion.py b/scepter/modules/model/network/diffusion/diffusion.py
index 5a23c56..4f8080a 100644
--- a/scepter/modules/model/network/diffusion/diffusion.py
+++ b/scepter/modules/model/network/diffusion/diffusion.py
@@ -54,7 +54,8 @@ class GaussianDiffusion(object):
guide_rescale=None,
clamp=None,
percentile=None,
- cat_uc=False):
+ cat_uc=False,
+ **kwargs):
"""
Apply one step of denoising from the posterior distribution q(x_s | x_t, x0).
Since x0 is not available, estimate the denoising results using the learned
@@ -79,7 +80,7 @@ class GaussianDiffusion(object):
# prediction
if guide_scale is None:
assert isinstance(model_kwargs, dict)
- out = model(xt, t=t, **model_kwargs)
+ out = model(xt, t=t, **model_kwargs, **kwargs)
else:
# classifier-free guidance (arXiv:2207.12598)
# model_kwargs[0]: conditional kwargs
@@ -87,7 +88,7 @@ class GaussianDiffusion(object):
assert isinstance(model_kwargs, list) and len(model_kwargs) == 2
if guide_scale == 1.:
- out = model(xt, t=t, **model_kwargs[0])
+ out = model(xt, t=t, **model_kwargs[0], **kwargs)
else:
if cat_uc:
@@ -111,11 +112,12 @@ class GaussianDiffusion(object):
all_model_kwargs[key], value)
all_out = model(xt.repeat(2, 1, 1, 1),
t=t.repeat(2),
- **all_model_kwargs)
+ **all_model_kwargs,
+ **kwargs)
y_out, u_out = all_out.chunk(2)
else:
- y_out = model(xt, t=t, **model_kwargs[0])
- u_out = model(xt, t=t, **model_kwargs[1])
+ y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
+ u_out = model(xt, t=t, **model_kwargs[1], **kwargs)
out = u_out + guide_scale * (y_out - u_out)
# rescale the output according to arXiv:2305.08891
@@ -262,7 +264,8 @@ class GaussianDiffusion(object):
guide_rescale,
clamp,
percentile,
- cat_uc=cat_uc)[-2]
+ cat_uc=cat_uc,
+ **kwargs)[-2]
# collect intermediate outputs
if return_intermediate == 'xt':
@@ -467,8 +470,8 @@ class GaussianDiffusion(object):
t = self._sigma_to_t(sigma).repeat(len(xt)).round().long()
x0 = self.denoise(xt * c_in, t, None, model, model_kwargs,
- guide_scale, guide_rescale, clamp,
- percentile)[-2]
+ guide_scale, guide_rescale, clamp, percentile,
+ **kwargs)[-2]
# collect intermediate outputs
if return_intermediate == 'xt':
intermediates.append(xt)
diff --git a/scepter/studio/inference/inference_ui/component_names.py b/scepter/studio/inference/inference_ui/component_names.py
index 1fb1de2..d25853f 100644
--- a/scepter/studio/inference/inference_ui/component_names.py
+++ b/scepter/studio/inference/inference_ui/component_names.py
@@ -83,7 +83,6 @@ class DiffusionUIName():
self.image_number = 'Images Number'
self.resolutions_height = 'Output Height'
self.resolutions_width = 'Output Width'
-
self.negative_prompt = 'Negative Prompt'
self.negative_prompt_placeholder = 'Type Prompt Here'
self.negative_prompt_description = 'Describing what you do not want to see.'
@@ -93,6 +92,12 @@ class DiffusionUIName():
self.discretization = 'Discretization'
self.random_seed = 'Use Random Seed'
self.seed = 'Used Seed'
+ self.example_block_name = 'Prompt Examples'
+ self.examples = [
+ 'dream dandelion', 'Mount Everest', 'a boy wearing a jacket',
+ 'Spring, Birds, Cawing, Branches',
+ 'Cyberpunk, Maiden, Heavy Machinery', 'A Frog'
+ ]
elif language == 'zh':
@@ -110,6 +115,12 @@ class DiffusionUIName():
self.discretization = '离散化'
self.random_seed = '使用随机种子'
self.seed = '使用的种子'
+ self.example_block_name = '提示词样例'
+ self.examples = [
+ 'dream dandelion', 'Mount Everest', 'a boy wearing a jacket',
+ 'Spring, Birds, Cawing, Branches',
+ 'Cyberpunk, Maiden, Heavy Machinery', 'A Frog'
+ ]
class MantraUIName():
@@ -127,6 +138,13 @@ class MantraUIName():
self.style_negative_template = 'Mantra Negative Prompt Template'
self.style_example = 'Mantra Results Example'
self.style_example_prompt = 'Mantra Example Prompt'
+ self.example_block_name = 'Examples'
+ self.examples = [
+ [['Adorable 3D Character'], 'a girl'],
+ [['Watercolor 2'], 'a single flower'],
+ [['Action Figure'],
+ 'a close up of a small rabbit wearing a hat and scarf']
+ ]
elif language == 'zh':
self.mantra_styles = '咒语风格'
@@ -140,6 +158,12 @@ class MantraUIName():
self.style_negative_template = '咒语负向提示模板'
self.style_example = '咒语示例图'
self.style_example_prompt = '咒语示例提示词'
+ self.example_block_name = '样例'
+ self.examples = [
+ [['可爱的3D角色'], 'a girl'], [['水彩'], 'a single flower'],
+ [['动作人偶'],
+ 'a close up of a small rabbit wearing a hat and scarf']
+ ]
class RefinerUIName():
@@ -179,6 +203,11 @@ class TunerUIName():
self.tuner_prompt_example = 'Prompt Example'
self.base_model = 'Base Model Name'
self.custom_tuner_model = 'Customized Model'
+ self.advance_block_name = 'Advance Setting'
+ self.tuner_scale = 'Tuner Scale'
+ self.example_block_name = 'Examples'
+ self.examples = [[['Pencil Sketch Drawing'], 'a girl in a jacket'],
+ [['Flat 2D Art'], 'a cat']]
elif language == 'zh':
self.tuner_model = '微调模型'
@@ -189,6 +218,11 @@ class TunerUIName():
self.tuner_prompt_example = '示例提示词'
self.base_model = '基础模型'
self.custom_tuner_model = '自定义模型'
+ self.advance_block_name = '高级设置'
+ self.tuner_scale = '微调强度'
+ self.example_block_name = '样例'
+ self.examples = [[['铅笔素描'], 'a girl in a jacket'],
+ [['扁平2D艺术'], 'a cat']]
class ControlUIName():
diff --git a/scepter/studio/inference/inference_ui/diffusion_ui.py b/scepter/studio/inference/inference_ui/diffusion_ui.py
index b1eb8dc..e4f0f34 100644
--- a/scepter/studio/inference/inference_ui/diffusion_ui.py
+++ b/scepter/studio/inference/inference_ui/diffusion_ui.py
@@ -68,6 +68,7 @@ class DiffusionUI(UIBase):
def create_ui(self, *args, **kwargs):
self.cur_paras = self.get_default(self.diffusion_paras,
self.default_input)
+ self.example_block = gr.Row(equal_height=True, visible=True)
with gr.Row(equal_height=True):
self.negative_prompt = gr.Textbox(
label=self.component_names.negative_prompt,
@@ -154,6 +155,12 @@ class DiffusionUI(UIBase):
self.refresh_seed = gr.Button(value=refresh_symbol)
def set_callbacks(self, model_manage_ui, **kwargs):
+ gallery_ui = kwargs.pop('gallery_ui')
+ with self.example_block:
+ gr.Examples(label=self.component_names.example_block_name,
+ examples=self.component_names.examples,
+ inputs=gallery_ui.prompt)
+
def random_checked(r):
value = -1
return (gr.Row(visible=not r), gr.Textbox(value=value))
diff --git a/scepter/studio/inference/inference_ui/gallery_ui.py b/scepter/studio/inference/inference_ui/gallery_ui.py
index 9460a9e..4313091 100644
--- a/scepter/studio/inference/inference_ui/gallery_ui.py
+++ b/scepter/studio/inference/inference_ui/gallery_ui.py
@@ -51,189 +51,203 @@ class GalleryUI(UIBase):
elem_id='generate_button',
visible=True)
+ def generate_gallery(self,
+ prompt,
+ mantra_state,
+ tuner_state,
+ control_state,
+ refine_state,
+ diffusion_model,
+ first_stage_model,
+ cond_stage_model,
+ refiner_cond_model,
+ refiner_diffusion_model,
+ tuner_model,
+ tuner_scale,
+ custom_tuner_model,
+ control_model,
+ control_scale,
+ crop_type,
+ control_cond_image,
+ negative_prompt,
+ prompt_prefix,
+ sample,
+ discretization,
+ output_height,
+ output_width,
+ image_number,
+ sample_steps,
+ guide_scale,
+ guide_rescale,
+ refine_strength,
+ refine_sampler,
+ refine_discretization,
+ refine_guide_scale,
+ refine_guide_rescale,
+ style_template,
+ style_negative_template,
+ image_seed,
+ show_jpeg_image=True):
+ current_pipeline = self.pipe_manager.get_pipeline_given_modules({
+ 'diffusion_model':
+ diffusion_model,
+ 'first_stage_model':
+ first_stage_model,
+ 'cond_stage_model':
+ cond_stage_model,
+ 'refiner_cond_model':
+ refiner_cond_model,
+ 'refiner_diffusion_model':
+ refiner_diffusion_model
+ })
+ now_pipeline = self.pipe_manager.model_level_info[diffusion_model][
+ 'pipeline'][0]
+ used_tuner_model = []
+ if not isinstance(tuner_model, list):
+ tuner_model = [tuner_model]
+ for tuner_m in tuner_model:
+ if tuner_m is None or tuner_m == '':
+ continue
+ if (now_pipeline in self.pipe_manager.model_level_info['tuners']
+ and tuner_m in self.pipe_manager.model_level_info['tuners']
+ [now_pipeline]):
+ tuner_m = self.pipe_manager.model_level_info['tuners'][
+ now_pipeline][tuner_m]['model_info']
+ used_tuner_model.append(tuner_m)
+ used_custom_tuner_model = []
+ if not isinstance(custom_tuner_model, list):
+ custom_tuner_model = [custom_tuner_model]
+ for tuner_m in custom_tuner_model:
+ if tuner_m is None or tuner_m == '':
+ continue
+ if (now_pipeline
+ in self.pipe_manager.model_level_info['customized_tuners']
+ and tuner_m in self.pipe_manager.
+ model_level_info['customized_tuners'][now_pipeline]):
+ tuner_m = self.pipe_manager.model_level_info[
+ 'customized_tuners'][now_pipeline][tuner_m]['model_info']
+ used_custom_tuner_model.append(tuner_m)
+
+ if (now_pipeline in self.pipe_manager.model_level_info['controllers']
+ and control_model in self.pipe_manager.
+ model_level_info['controllers'][now_pipeline]):
+ control_model = self.pipe_manager.model_level_info['controllers'][
+ now_pipeline][control_model]['model_info']
+
+ prompt_rephrased = style_template.replace(
+ '{prompt}',
+ prompt) if not style_template == '' and mantra_state else prompt
+ prompt_rephrased = f'{prompt_prefix}{prompt_rephrased}' if not prompt_prefix == '' else prompt_rephrased
+ negative_prompt_rephrased = negative_prompt + style_negative_template if mantra_state else negative_prompt
+ pipeline_input = {
+ 'prompt': prompt_rephrased,
+ 'negative_prompt': negative_prompt_rephrased,
+ 'sample': sample,
+ 'sample_steps': sample_steps,
+ 'discretization': discretization,
+ 'original_size_as_tuple': [int(output_height),
+ int(output_width)],
+ 'target_size_as_tuple': [int(output_height),
+ int(output_width)],
+ 'crop_coords_top_left': [0, 0],
+ 'guide_scale': guide_scale,
+ 'guide_rescale': guide_rescale,
+ }
+ if refine_state:
+ pipeline_input['refine_sampler'] = refine_sampler
+ pipeline_input['refine_discretization'] = refine_discretization
+ pipeline_input['refine_guide_scale'] = refine_guide_scale
+ pipeline_input['refine_guide_rescale'] = refine_guide_rescale
+ else:
+ refine_strength = 0
+ results = current_pipeline(
+ pipeline_input,
+ num_samples=image_number,
+ intermediate_callback=None,
+ refine_strength=refine_strength,
+ img_to_img_strength=0,
+ tuner_model=used_tuner_model +
+ used_custom_tuner_model if tuner_state else None,
+ tuner_scale=tuner_scale if tuner_state or control_state else None,
+ control_model=control_model if control_state else None,
+ control_scale=control_scale
+ if tuner_state or control_state else None,
+ control_cond_image=control_cond_image if control_state else None,
+ crop_type=crop_type if control_state else None,
+ seed=int(image_seed))
+ images = []
+ before_images = []
+ if 'images' in results:
+ images_tensor = results['images'] * 255
+ images = [
+ Image.fromarray(images_tensor[idx].permute(
+ 1, 2, 0).cpu().numpy().astype(np.uint8))
+ for idx in range(images_tensor.shape[0])
+ ]
+ if 'before_refine_images' in results and results[
+ 'before_refine_images'] is not None:
+ before_refine_images_tensor = results['before_refine_images'] * 255
+ before_images = [
+ Image.fromarray(before_refine_images_tensor[idx].permute(
+ 1, 2, 0).cpu().numpy().astype(np.uint8))
+ for idx in range(before_refine_images_tensor.shape[0])
+ ]
+ if 'seed' in results:
+ print(results['seed'])
+ print(images, before_images)
+ if show_jpeg_image:
+ save_list = []
+ for i, img in enumerate(images):
+ save_image = os.path.join(self.cfg.WORK_DIR,
+ f'cur_gallery_{i}.jpg')
+ img.save(save_image)
+ save_list.append(save_image)
+ images = save_list
+ return (
+ gr.Column(visible=len(before_images) > 0),
+ before_images,
+ images,
+ )
+
+ def generate_image(self, *args, **kwargs):
+ gallery_result = self.generate_gallery(*args, **kwargs)
+ before_refine_panel, before_refine_gallery, output_gallery = gallery_result
+ return (before_refine_panel, before_refine_gallery, output_gallery[0])
+
def set_callbacks(self, inference_ui, model_manage_ui, diffusion_ui,
mantra_ui, tuner_ui, refiner_ui, control_ui, **kwargs):
- def generate_image(mantra_state,
- tuner_state,
- control_state,
- refine_state,
- diffusion_model,
- first_stage_model,
- cond_stage_model,
- refiner_cond_model,
- refiner_diffusion_model,
- tuner_model,
- custom_tuner_model,
- control_model,
- crop_type,
- control_cond_image,
- prompt,
- negative_prompt,
- prompt_prefix,
- sample,
- discretization,
- output_height,
- output_width,
- image_number,
- sample_steps,
- guide_scale,
- guide_rescale,
- refine_strength,
- refine_sampler,
- refine_discretization,
- refine_guide_scale,
- refine_guide_rescale,
- style_template,
- style_negative_template,
- image_seed,
- show_jpeg_image=True):
- current_pipeline = self.pipe_manager.get_pipeline_given_modules({
- 'diffusion_model':
- diffusion_model,
- 'first_stage_model':
- first_stage_model,
- 'cond_stage_model':
- cond_stage_model,
- 'refiner_cond_model':
- refiner_cond_model,
- 'refiner_diffusion_model':
- refiner_diffusion_model
- })
- now_pipeline = self.pipe_manager.model_level_info[diffusion_model][
- 'pipeline'][0]
- used_tuner_model = []
- if not isinstance(tuner_model, list):
- tuner_model = [tuner_model]
- for tuner_m in tuner_model:
- if tuner_m is None or tuner_m == '':
- continue
- if (now_pipeline
- in self.pipe_manager.model_level_info['tuners']
- and tuner_m in self.pipe_manager.
- model_level_info['tuners'][now_pipeline]):
- tuner_m = self.pipe_manager.model_level_info['tuners'][
- now_pipeline][tuner_m]['model_info']
- used_tuner_model.append(tuner_m)
- used_custom_tuner_model = []
- if not isinstance(custom_tuner_model, list):
- custom_tuner_model = [custom_tuner_model]
- for tuner_m in custom_tuner_model:
- if tuner_m is None or tuner_m == '':
- continue
- if (now_pipeline in
- self.pipe_manager.model_level_info['customized_tuners']
- and tuner_m in self.pipe_manager.
- model_level_info['customized_tuners'][now_pipeline]):
- tuner_m = self.pipe_manager.model_level_info[
- 'customized_tuners'][now_pipeline][tuner_m][
- 'model_info']
- used_custom_tuner_model.append(tuner_m)
- if (now_pipeline
- in self.pipe_manager.model_level_info['controllers']
- and control_model in self.pipe_manager.
- model_level_info['controllers'][now_pipeline]):
- control_model = self.pipe_manager.model_level_info[
- 'controllers'][now_pipeline][control_model]['model_info']
+ self.gen_inputs = [
+ self.prompt, mantra_ui.state, tuner_ui.state, control_ui.state,
+ refiner_ui.state, model_manage_ui.diffusion_model,
+ model_manage_ui.first_stage_model,
+ model_manage_ui.cond_stage_model, refiner_ui.refiner_cond_model,
+ refiner_ui.refiner_diffusion_model, tuner_ui.tuner_model,
+ tuner_ui.tuner_scale, tuner_ui.custom_tuner_model,
+ control_ui.control_model, control_ui.control_scale,
+ control_ui.crop_type, control_ui.cond_image,
+ diffusion_ui.negative_prompt, diffusion_ui.prompt_prefix,
+ diffusion_ui.sampler, diffusion_ui.discretization,
+ diffusion_ui.output_height, diffusion_ui.output_width,
+ diffusion_ui.image_number, diffusion_ui.sample_steps,
+ diffusion_ui.guide_scale, diffusion_ui.guide_rescale,
+ refiner_ui.refine_strength, refiner_ui.refine_sampler,
+ refiner_ui.refine_discretization, refiner_ui.refine_guide_scale,
+ refiner_ui.refine_guide_rescale, mantra_ui.style_template,
+ mantra_ui.style_negative_template, diffusion_ui.image_seed
+ ]
- prompt_rephrased = style_template.replace(
- '{prompt}', prompt
- ) if not style_template == '' and mantra_state else prompt
- prompt_rephrased = f'{prompt_prefix}{prompt_rephrased}' if not prompt_prefix == '' else prompt_rephrased
- negative_prompt_rephrased = negative_prompt + style_negative_template if mantra_state else negative_prompt
- pipeline_input = {
- 'prompt': prompt_rephrased,
- 'negative_prompt': negative_prompt_rephrased,
- 'sample': sample,
- 'sample_steps': sample_steps,
- 'discretization': discretization,
- 'original_size_as_tuple':
- [int(output_height), int(output_width)],
- 'target_size_as_tuple':
- [int(output_height), int(output_width)],
- 'crop_coords_top_left': [0, 0],
- 'guide_scale': guide_scale,
- 'guide_rescale': guide_rescale,
- }
- if refine_state:
- pipeline_input['refine_sampler'] = refine_sampler
- pipeline_input['refine_discretization'] = refine_discretization
- pipeline_input['refine_guide_scale'] = refine_guide_scale
- pipeline_input['refine_guide_rescale'] = refine_guide_rescale
- else:
- refine_strength = 0
- results = current_pipeline(
- pipeline_input,
- num_samples=image_number,
- intermediate_callback=None,
- refine_strength=refine_strength,
- img_to_img_strength=0,
- tuner_model=used_tuner_model +
- used_custom_tuner_model if tuner_state else None,
- control_model=control_model if control_state else None,
- control_cond_image=control_cond_image
- if control_state else None,
- crop_type=crop_type if control_state else None,
- seed=int(image_seed))
- images = []
- before_images = []
- if 'images' in results:
- images_tensor = results['images'] * 255
- images = [
- Image.fromarray(images_tensor[idx].permute(
- 1, 2, 0).cpu().numpy().astype(np.uint8))
- for idx in range(images_tensor.shape[0])
- ]
- if 'before_refine_images' in results and results[
- 'before_refine_images'] is not None:
- before_refine_images_tensor = results[
- 'before_refine_images'] * 255
- before_images = [
- Image.fromarray(before_refine_images_tensor[idx].permute(
- 1, 2, 0).cpu().numpy().astype(np.uint8))
- for idx in range(before_refine_images_tensor.shape[0])
- ]
- if 'seed' in results:
- print(results['seed'])
- print(images, before_images)
- if show_jpeg_image:
- save_list = []
- for i, img in enumerate(images):
- save_image = os.path.join(self.cfg.WORK_DIR,
- f'cur_gallery_{i}.jpg')
- img.save(save_image)
- save_list.append(save_image)
- images = save_list
- return (
- gr.Column(visible=len(before_images) > 0),
- before_images,
- images,
- )
+ self.gen_outputs = [
+ self.before_refine_panel, self.before_refine_gallery,
+ self.output_gallery
+ ]
- self.generate_button.click(
- generate_image,
- inputs=[
- mantra_ui.state, tuner_ui.state, control_ui.state,
- refiner_ui.state, model_manage_ui.diffusion_model,
- model_manage_ui.first_stage_model,
- model_manage_ui.cond_stage_model,
- refiner_ui.refiner_cond_model,
- refiner_ui.refiner_diffusion_model, tuner_ui.tuner_model,
- tuner_ui.custom_tuner_model, control_ui.control_model,
- control_ui.crop_type, control_ui.cond_image, self.prompt,
- diffusion_ui.negative_prompt, diffusion_ui.prompt_prefix,
- diffusion_ui.sampler, diffusion_ui.discretization,
- diffusion_ui.output_height, diffusion_ui.output_width,
- diffusion_ui.image_number, diffusion_ui.sample_steps,
- diffusion_ui.guide_scale, diffusion_ui.guide_rescale,
- refiner_ui.refine_strength, refiner_ui.refine_sampler,
- refiner_ui.refine_discretization,
- refiner_ui.refine_guide_scale, refiner_ui.refine_guide_rescale,
- mantra_ui.style_template, mantra_ui.style_negative_template,
- diffusion_ui.image_seed
- ],
- outputs=[
- self.before_refine_panel, self.before_refine_gallery,
- self.output_gallery
- ],
- queue=True)
+ self.generate_button.click(self.generate_gallery,
+ inputs=self.gen_inputs,
+ outputs=self.gen_outputs,
+ queue=True)
+
+ self.prompt.submit(self.generate_gallery,
+ inputs=self.gen_inputs,
+ outputs=self.gen_outputs,
+ queue=True)
diff --git a/scepter/studio/inference/inference_ui/mantra_ui.py b/scepter/studio/inference/inference_ui/mantra_ui.py
index 22143f0..1be6c42 100644
--- a/scepter/studio/inference/inference_ui/mantra_ui.py
+++ b/scepter/studio/inference/inference_ui/mantra_ui.py
@@ -47,64 +47,75 @@ class MantraUI(UIBase):
def create_ui(self, *args, **kwargs):
self.state = gr.State(value=False)
- with gr.Row(equal_height=True, visible=False) as self.tab:
- with gr.Column(scale=1):
- with gr.Group(visible=True):
- with gr.Row(equal_height=True):
- self.style = gr.Dropdown(
- label=self.component_names.mantra_styles,
- choices=self.all_styles[self.default_pipeline],
- value=None,
- multiselect=True,
- interactive=True)
- with gr.Row(equal_height=True):
- with gr.Column(scale=1):
- self.style_name = gr.Text(
+ with gr.Column(visible=False) as self.tab:
+ with gr.Row():
+ with gr.Column(scale=1):
+ with gr.Group(visible=True):
+ with gr.Row(equal_height=True):
+ self.style = gr.Dropdown(
+ label=self.component_names.mantra_styles,
+ choices=self.all_styles[self.default_pipeline],
+ value=None,
+ multiselect=True,
+ interactive=True)
+ with gr.Row(equal_height=True):
+ with gr.Column(scale=1):
+ self.style_name = gr.Text(
+ value='',
+ label=self.component_names.style_name)
+ with gr.Column(scale=1):
+ self.style_source = gr.Text(
+ value='',
+ label=self.component_names.style_source)
+ with gr.Column(scale=1):
+ self.style_desc = gr.Text(
+ value='',
+ label=self.component_names.style_desc)
+ with gr.Row(equal_height=True):
+ self.style_prompt = gr.Text(
value='',
- label=self.component_names.style_name)
- with gr.Column(scale=1):
- self.style_source = gr.Text(
+ label=self.component_names.style_prompt,
+ lines=4)
+ with gr.Row(equal_height=True):
+ self.style_negative_prompt = gr.Text(
value='',
- label=self.component_names.style_source)
- with gr.Column(scale=1):
- self.style_desc = gr.Text(
+ label=self.component_names.
+ style_negative_prompt,
+ lines=4)
+ with gr.Column(scale=1):
+ with gr.Group(visible=True):
+ with gr.Row(equal_height=True):
+ self.style_template = gr.Text(
value='',
- label=self.component_names.style_desc)
- with gr.Row(equal_height=True):
- self.style_prompt = gr.Text(
- value='',
- label=self.component_names.style_prompt,
- lines=4)
- with gr.Row(equal_height=True):
- self.style_negative_prompt = gr.Text(
- value='',
- label=self.component_names.style_negative_prompt,
- lines=4)
- with gr.Column(scale=1):
- with gr.Group(visible=True):
- with gr.Row(equal_height=True):
- self.style_template = gr.Text(
- value='',
- label=self.component_names.style_template,
- lines=2)
- with gr.Row(equal_height=True):
- self.style_negative_template = gr.Text(
- value='',
- label=self.component_names.style_negative_template,
- lines=2)
- with gr.Row(equal_height=True):
- self.style_example = gr.Image(
- label=self.component_names.style_example,
- source='upload',
- value=None,
- interactive=False)
- with gr.Row(equal_height=True):
- self.style_example_prompt = gr.Text(
- value='',
- label=self.component_names.style_example_prompt,
- lines=2)
+ label=self.component_names.style_template,
+ lines=2)
+ with gr.Row(equal_height=True):
+ self.style_negative_template = gr.Text(
+ value='',
+ label=self.component_names.
+ style_negative_template,
+ lines=2)
+ with gr.Row(equal_height=True):
+ self.style_example = gr.Image(
+ label=self.component_names.style_example,
+ source='upload',
+ value=None,
+ interactive=False)
+ with gr.Row(equal_height=True):
+ self.style_example_prompt = gr.Text(
+ value='',
+ label=self.component_names.
+ style_example_prompt,
+ lines=2)
+ self.example_block = gr.Accordion(
+ label=self.component_names.example_block_name, open=True)
def set_callbacks(self, model_manage_ui, **kwargs):
+ gallery_ui = kwargs.pop('gallery_ui')
+ with self.example_block:
+ gr.Examples(examples=self.component_names.examples,
+ inputs=[self.style, gallery_ui.prompt])
+
def change_style(style, diffusion_model):
style_template = ''
style_negative_template = []
diff --git a/scepter/studio/inference/inference_ui/tuner_ui.py b/scepter/studio/inference/inference_ui/tuner_ui.py
index 4e73f93..4ca4f06 100644
--- a/scepter/studio/inference/inference_ui/tuner_ui.py
+++ b/scepter/studio/inference/inference_ui/tuner_ui.py
@@ -47,53 +47,74 @@ class TunerUI(UIBase):
def create_ui(self, *args, **kwargs):
self.state = gr.State(value=False)
- with gr.Row(equal_height=True, visible=False) as self.tab:
- with gr.Column(variant='panel', scale=1, min_width=0):
- with gr.Group(visible=True):
- with gr.Row(equal_height=True):
- with gr.Column(scale=1):
- self.tuner_model = gr.Dropdown(
- label=self.component_names.tuner_model,
- choices=self.tunner_choices,
+ with gr.Column(visible=False) as self.tab:
+ with gr.Row():
+ with gr.Column(variant='panel', scale=1, min_width=0):
+ with gr.Group(visible=True):
+ with gr.Row(equal_height=True):
+ with gr.Column(scale=1):
+ self.tuner_model = gr.Dropdown(
+ label=self.component_names.tuner_model,
+ choices=self.tunner_choices,
+ value=None,
+ multiselect=True,
+ interactive=True)
+ with gr.Column(scale=1):
+ self.custom_tuner_model = gr.Dropdown(
+ label=self.component_names.
+ custom_tuner_model,
+ choices=[],
+ value=None,
+ multiselect=True,
+ interactive=True)
+ with gr.Row(equal_height=True):
+ with gr.Column(scale=1):
+ self.tuner_type = gr.Text(
+ value='',
+ label=self.component_names.tuner_type)
+ with gr.Column(scale=1):
+ self.base_model = gr.Text(
+ value='',
+ label=self.component_names.base_model)
+ with gr.Column(scale=1):
+ self.tuner_desc = gr.Text(
+ value='',
+ label=self.component_names.tuner_desc,
+ lines=4)
+ with gr.Column(variant='panel', scale=1, min_width=0):
+ with gr.Group(visible=True):
+ with gr.Row(equal_height=True):
+ self.tuner_example = gr.Image(
+ label=self.component_names.tuner_example,
+ source='upload',
value=None,
- multiselect=True,
- interactive=True)
- with gr.Column(scale=1):
- self.custom_tuner_model = gr.Dropdown(
- label=self.component_names.custom_tuner_model,
- choices=[],
- value=None,
- multiselect=True,
- interactive=True)
- with gr.Row(equal_height=True):
- with gr.Column(scale=1):
- self.tuner_type = gr.Text(
+ interactive=False)
+ with gr.Row(equal_height=True):
+ self.tuner_prompt_example = gr.Text(
value='',
- label=self.component_names.tuner_type)
- with gr.Column(scale=1):
- self.base_model = gr.Text(
- value='',
- label=self.component_names.base_model)
- with gr.Column(scale=1):
- self.tuner_desc = gr.Text(
- value='',
- label=self.component_names.tuner_desc,
- lines=4)
- with gr.Column(variant='panel', scale=1, min_width=0):
- with gr.Group(visible=True):
- with gr.Row(equal_height=True):
- self.tuner_example = gr.Image(
- label=self.component_names.tuner_example,
- source='upload',
- value=None,
- interactive=False)
- with gr.Row(equal_height=True):
- self.tuner_prompt_example = gr.Text(
- value='',
- label=self.component_names.tuner_prompt_example,
- lines=2)
+ label=self.component_names.
+ tuner_prompt_example,
+ lines=2)
+
+ with gr.Accordion(label=self.component_names.advance_block_name,
+ open=False):
+ self.tuner_scale = gr.Slider(
+ label=self.component_names.tuner_scale,
+ minimum=0.0,
+ maximum=1.0,
+ step=0.05,
+ value=1.0,
+ interactive=True)
+
+ self.example_block = gr.Accordion(
+ label=self.component_names.example_block_name, open=True)
def set_callbacks(self, model_manage_ui, **kwargs):
+ gallery_ui = kwargs.pop('gallery_ui')
+ with self.example_block:
+ gr.Examples(examples=self.component_names.examples,
+ inputs=[self.tuner_model, gallery_ui.prompt])
+
def tuner_model_change(tuner_model, diffusion_model):
diffusion_model_info = self.pipe_manager.model_level_info[
diffusion_model]
diff --git a/scepter/studio/preprocess/caption_editor_ui/dataset_gallery_ui.py b/scepter/studio/preprocess/caption_editor_ui/dataset_gallery_ui.py
index 01f83ea..df8a730 100644
--- a/scepter/studio/preprocess/caption_editor_ui/dataset_gallery_ui.py
+++ b/scepter/studio/preprocess/caption_editor_ui/dataset_gallery_ui.py
@@ -5,8 +5,8 @@ from __future__ import annotations
import os.path
import gradio as gr
-
import imagehash
+
from scepter.modules.utils.file_system import FS
from scepter.studio.preprocess.caption_editor_ui.component_names import \
DatasetGalleryUIName