stable-diffusion-webui
121 строка · 4.7 Кб
1import json
2from contextlib import closing
3
4import modules.scripts
5from modules import processing, infotext_utils
6from modules.infotext_utils import create_override_settings_dict, parse_generation_parameters
7from modules.shared import opts
8import modules.shared as shared
9from modules.ui import plaintext_to_html
10from PIL import Image
11import gradio as gr
12
13
14def txt2img_create_processing(id_task: str, request: gr.Request, prompt: str, negative_prompt: str, prompt_styles, steps: int, sampler_name: str, n_iter: int, batch_size: int, cfg_scale: float, height: int, width: int, enable_hr: bool, denoising_strength: float, hr_scale: float, hr_upscaler: str, hr_second_pass_steps: int, hr_resize_x: int, hr_resize_y: int, hr_checkpoint_name: str, hr_sampler_name: str, hr_prompt: str, hr_negative_prompt, override_settings_texts, *args, force_enable_hr=False):
15override_settings = create_override_settings_dict(override_settings_texts)
16
17if force_enable_hr:
18enable_hr = True
19
20p = processing.StableDiffusionProcessingTxt2Img(
21sd_model=shared.sd_model,
22outpath_samples=opts.outdir_samples or opts.outdir_txt2img_samples,
23outpath_grids=opts.outdir_grids or opts.outdir_txt2img_grids,
24prompt=prompt,
25styles=prompt_styles,
26negative_prompt=negative_prompt,
27sampler_name=sampler_name,
28batch_size=batch_size,
29n_iter=n_iter,
30steps=steps,
31cfg_scale=cfg_scale,
32width=width,
33height=height,
34enable_hr=enable_hr,
35denoising_strength=denoising_strength,
36hr_scale=hr_scale,
37hr_upscaler=hr_upscaler,
38hr_second_pass_steps=hr_second_pass_steps,
39hr_resize_x=hr_resize_x,
40hr_resize_y=hr_resize_y,
41hr_checkpoint_name=None if hr_checkpoint_name == 'Use same checkpoint' else hr_checkpoint_name,
42hr_sampler_name=None if hr_sampler_name == 'Use same sampler' else hr_sampler_name,
43hr_prompt=hr_prompt,
44hr_negative_prompt=hr_negative_prompt,
45override_settings=override_settings,
46)
47
48p.scripts = modules.scripts.scripts_txt2img
49p.script_args = args
50
51p.user = request.username
52
53if shared.opts.enable_console_prompts:
54print(f"\ntxt2img: {prompt}", file=shared.progress_print_out)
55
56return p
57
58
59def txt2img_upscale(id_task: str, request: gr.Request, gallery, gallery_index, generation_info, *args):
60assert len(gallery) > 0, 'No image to upscale'
61assert 0 <= gallery_index < len(gallery), f'Bad image index: {gallery_index}'
62
63p = txt2img_create_processing(id_task, request, *args, force_enable_hr=True)
64p.batch_size = 1
65p.n_iter = 1
66# txt2img_upscale attribute that signifies this is called by txt2img_upscale
67p.txt2img_upscale = True
68
69geninfo = json.loads(generation_info)
70
71image_info = gallery[gallery_index] if 0 <= gallery_index < len(gallery) else gallery[0]
72p.firstpass_image = infotext_utils.image_from_url_text(image_info)
73
74parameters = parse_generation_parameters(geninfo.get('infotexts')[gallery_index], [])
75p.seed = parameters.get('Seed', -1)
76p.subseed = parameters.get('Variation seed', -1)
77
78p.override_settings['save_images_before_highres_fix'] = False
79
80with closing(p):
81processed = modules.scripts.scripts_txt2img.run(p, *p.script_args)
82
83if processed is None:
84processed = processing.process_images(p)
85
86shared.total_tqdm.clear()
87
88new_gallery = []
89for i, image in enumerate(gallery):
90if i == gallery_index:
91geninfo["infotexts"][gallery_index: gallery_index+1] = processed.infotexts
92new_gallery.extend(processed.images)
93else:
94fake_image = Image.new(mode="RGB", size=(1, 1))
95fake_image.already_saved_as = image["name"].rsplit('?', 1)[0]
96new_gallery.append(fake_image)
97
98geninfo["infotexts"][gallery_index] = processed.info
99
100return new_gallery, json.dumps(geninfo), plaintext_to_html(processed.info), plaintext_to_html(processed.comments, classname="comments")
101
102
103def txt2img(id_task: str, request: gr.Request, *args):
104p = txt2img_create_processing(id_task, request, *args)
105
106with closing(p):
107processed = modules.scripts.scripts_txt2img.run(p, *p.script_args)
108
109if processed is None:
110processed = processing.process_images(p)
111
112shared.total_tqdm.clear()
113
114generation_info_js = processed.js()
115if opts.samples_log_stdout:
116print(generation_info_js)
117
118if opts.do_not_show_images:
119processed.images = []
120
121return processed.images, generation_info_js, plaintext_to_html(processed.info), plaintext_to_html(processed.comments, classname="comments")
122