Merge pull request #4547 from awsr/typing

Typing and related
This commit is contained in:
Vladimir Mandic
2026-01-14 08:38:21 +01:00
committed by GitHub
10 changed files with 70 additions and 62 deletions
+2 -2
View File
@@ -1,6 +1,6 @@
import os
import time
from typing import Any, Dict
from typing import Any
from fastapi import Request, Depends
from fastapi.exceptions import HTTPException
from fastapi.responses import FileResponse
@@ -95,7 +95,7 @@ def get_config():
del options['sd_lora']
return options
def set_config(req: Dict[str, Any]):
def set_config(req: dict[str, Any]):
updated = []
for k, v in req.items():
updated.append({ k: shared.opts.set(k, v) })
+11 -10
View File
@@ -1,3 +1,4 @@
from __future__ import annotations
import base64
import io
import os
@@ -8,9 +9,9 @@ from modules.infotext import parse, mapping, quote, unquote # pylint: disable=un
type_of_gr_update = type(gr.update())
paste_fields = {}
paste_fields: dict[str, dict] = {}
field_names = {}
registered_param_bindings = []
registered_param_bindings: list[ParamBinding] = []
debug = shared.log.trace if os.environ.get('SD_PASTE_DEBUG', None) is not None else lambda *args, **kwargs: None
debug('Trace: PASTE')
parse_generation_parameters = parse # compatibility
@@ -18,7 +19,7 @@ infotext_to_setting_name_mapping = mapping # compatibility
class ParamBinding:
def __init__(self, paste_button, tabname, source_text_component=None, source_image_component=None, source_tabname=None, override_settings_component=None, paste_field_names=None):
def __init__(self, paste_button, tabname: str, source_text_component=None, source_image_component=None, source_tabname=None, override_settings_component=None, paste_field_names=None):
self.paste_button = paste_button
self.tabname = tabname
self.source_text_component = source_text_component
@@ -60,7 +61,7 @@ def image_from_url_text(filedata):
if len(filedata) == 0:
return None
filedata = filedata[0]
if type(filedata) == dict:
if not isinstance(filedata, str):
shared.log.warning('Incorrect filedata received')
return None
if filedata.startswith("data:image/png;base64,"):
@@ -71,13 +72,13 @@ def image_from_url_text(filedata):
filedata = filedata[len("data:image/jpeg;base64,"):]
if filedata.startswith("data:image/jxl;base64,"):
filedata = filedata[len("data:image/jxl;base64,"):]
filedata = base64.decodebytes(filedata.encode('utf-8'))
image = Image.open(io.BytesIO(filedata))
filebytes = base64.decodebytes(filedata.encode('utf-8'))
image = Image.open(io.BytesIO(filebytes))
images.read_info_from_image(image)
return image
def add_paste_fields(tabname, init_img, fields, override_settings_component=None):
def add_paste_fields(tabname: str, init_img: gr.Image | gr.HTML | None, fields: list[tuple[gr.components.Component, str]] | None, override_settings_component=None):
paste_fields[tabname] = {"init_img": init_img, "fields": fields, "override_settings_component": override_settings_component}
try:
field_names[tabname] = [f[1] for f in fields if f[1] is not None and not callable(f[1])] if fields is not None else [] # tuple (component, label)
@@ -108,7 +109,7 @@ def get_all_fields():
return all_fields
def create_buttons(tabs_list):
def create_buttons(tabs_list: list[str]) -> dict[str, gr.Button]:
buttons = {}
for tab in tabs_list:
name = tab
@@ -128,7 +129,7 @@ def create_buttons(tabs_list):
return buttons
def should_skip(param):
def should_skip(param: str):
skip_params = [p.strip().lower() for p in shared.opts.disable_apply_params.split(",")]
if not shared.opts.clip_skip_enabled:
skip_params += ['clip skip']
@@ -149,7 +150,7 @@ def connect_paste_params_buttons():
if binding.tabname not in paste_fields:
debug(f"Not not registered: tab={binding.tabname}")
continue
fields = paste_fields[binding.tabname]["fields"]
fields: list[tuple[gr.components.Component, str]] = paste_fields[binding.tabname]["fields"]
destination_image_component = paste_fields[binding.tabname]["init_img"]
if binding.source_image_component:
+2 -2
View File
@@ -311,7 +311,7 @@ def parse_novelai_metadata(data: dict):
return geninfo
def read_info_from_image(image: Image, watermark: bool = False):
def read_info_from_image(image: Image.Image, watermark: bool = False):
if image is None:
return '', {}
if isinstance(image, str):
@@ -419,7 +419,7 @@ def draw_overlay(im, text: str = '', y_offset: int = 0):
return im
def set_watermark(image, wm_text: str = None, wm_image: Image.Image = None):
def set_watermark(image, wm_text: str | None = None, wm_image: Image.Image | None = None):
if shared.opts.image_watermark_position != 'none' and wm_image is not None: # visible watermark
if isinstance(wm_image, str):
try:
+4 -3
View File
@@ -25,7 +25,7 @@ def readfile(filename: str, silent: bool = False, lock: bool = False, *, as_type
if lock and locking_available:
try:
lock_file = fasteners.InterProcessReaderWriterLock(f"{filename}.lock")
lock_file.logger.disabled = True
lock_file.logger.disabled = True # type: ignore - False positive. Bad typing in Fasteners.
locked = lock_file.acquire_read_lock(blocking=True, timeout=3)
except Exception as err:
lock_file = None
@@ -105,7 +105,7 @@ def writefile(data, filename, mode='w', silent=False, atomic=False):
try:
if locking_available:
lock_file = fasteners.InterProcessReaderWriterLock(f"{filename}.lock") if locking_available else None
lock_file.logger.disabled = True
lock_file.logger.disabled = True # type: ignore - False positive. Bad typing in Fasteners.
locked = lock_file.acquire_write_lock(blocking=True, timeout=3) if lock_file is not None else False
except Exception as err:
locking_available = False
@@ -124,7 +124,8 @@ def writefile(data, filename, mode='w', silent=False, atomic=False):
file.write(output)
t1 = time.time()
if not silent:
log.debug(f'Save: file="{filename}" json={len(data)} bytes={len(output)} time={t1-t0:.3f}')
datalength = len(data) if isinstance(data, (dict, list)) else (len(data.__dict__))
log.debug(f'Save: file="{filename}" json={datalength} bytes={len(output)} time={t1-t0:.3f}')
except Exception as err:
log.error(f'Save failed: file="{filename}" {err}')
try:
+1 -1
View File
@@ -21,7 +21,7 @@ def options_section(section_identifier: tuple[str, str], options_dict: dict[str,
class OptionInfo:
def __init__(
self,
default: Any | None = None,
default: Any = None,
label="",
component: type[Component] | type[DropdownEditable] | None = None,
component_args: dict | Callable[..., dict] | None = None,
+27 -28
View File
@@ -12,56 +12,54 @@ from installer import log
if TYPE_CHECKING:
from collections.abc import Callable
from modules.options import OptionInfo
from typing import Any
cmd_opts = cmd_args.parse_args()
compatibility_opts = ['clip_skip', 'uni_pc_lower_order_final', 'uni_pc_order']
class Options():
data = None
data_labels = None
data_labels: dict[str, OptionInfo | LegacyOption]
data: dict[str, Any]
typemap = {int: float}
debug = os.environ.get('SD_CONFIG_DEBUG', None) is not None
def __init__(self, options_templates: dict[str, OptionInfo | LegacyOption] = {}, restricted_opts: set[str] | None = None, *, filename = ''):
if restricted_opts is None:
restricted_opts = set()
super().__setattr__('data_labels', options_templates)
super().__setattr__('data', {k: v.default for k, v in options_templates.items()})
self.filename: str = filename or cmd_opts.config
self.data_labels = options_templates
self.restricted_opts = restricted_opts
self.data = {k: v.default for k, v in self.data_labels.items()}
self.legacy = [k for k, v in self.data_labels.items() if isinstance(v, LegacyOption)]
self.legacy = [k for k, v in options_templates.items() if isinstance(v, LegacyOption)]
self.load()
def __setattr__(self, key, value): # pylint: disable=inconsistent-return-statements
if self.data is not None:
if key in self.data or key in self.data_labels:
if cmd_opts.freeze:
log.warning(f'Settings are frozen: {key}')
return
if cmd_opts.hide_ui_dir_config and key in self.restricted_opts:
log.warning(f'Settings key is restricted: {key}')
return
if self.debug:
log.trace(f'Settings set: {key}={value}')
if key in self.legacy:
log.warning(f'Settings set: {key}={value} legacy')
self.data[key] = value
if key in self.data or key in self.data_labels:
if cmd_opts.freeze:
log.warning(f'Settings are frozen: {key}')
return
if cmd_opts.hide_ui_dir_config and key in self.restricted_opts:
log.warning(f'Settings key is restricted: {key}')
return
if self.debug:
log.trace(f'Settings set: {key}={value}')
if key in self.legacy:
log.warning(f'Settings set: {key}={value} legacy')
self.data[key] = value
return
return super(Options, self).__setattr__(key, value) # pylint: disable=super-with-arguments
def get(self, item):
if self.data is not None:
if item in self.data:
return self.data[item]
if item in self.data:
return self.data[item]
if item in self.data_labels:
return self.data_labels[item].default
return super(Options, self).__getattribute__(item) # pylint: disable=super-with-arguments
def __getattr__(self, item):
if self.data is not None:
if item in self.data:
return self.data[item]
if item in self.data:
return self.data[item]
if item in self.data_labels:
return self.data_labels[item].default
return super(Options, self).__getattribute__(item) # pylint: disable=super-with-arguments
@@ -81,9 +79,10 @@ class Options():
setattr(self, key, value)
except RuntimeError:
return False
if self.data_labels[key].onchange is not None:
func = self.data_labels[key].onchange
if func is not None:
try:
self.data_labels[key].onchange()
func()
except Exception as err:
log.error(f'Error in onchange callback: {key} {value} {err}')
errors.display(err, 'Error in onchange callback')
@@ -169,10 +168,10 @@ class Options():
return
self.data = readfile(filename, lock=True, as_type="dict")
if self.data.get('quicksettings') is not None and self.data.get('quicksettings_list') is None:
self.data['quicksettings_list'] = [i.strip() for i in self.data.get('quicksettings').split(',')]
self.data['quicksettings_list'] = [i.strip() for i in self.data.get('quicksettings', '').split(',')]
unknown_settings = []
for k, v in self.data.items():
info: OptionInfo | None = self.data_labels.get(k, None)
info = self.data_labels.get(k, None)
if info is not None:
if not info.validate(k, v):
self.data[k] = info.default
+2 -3
View File
@@ -1,4 +1,3 @@
from typing import List
import io
import time
import json
@@ -47,7 +46,7 @@ dtypes = {
}
def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_type: str = None) -> Image.Image:
def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_type: str | None = None):
from modules import devices, shared, errors, modelloader
tensors = []
content = 0
@@ -128,7 +127,7 @@ def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_
return tensors
def remote_encode(images: List[Image.Image], model_type: str = None) -> torch.Tensor:
def remote_encode(images: list[Image.Image], model_type: str | None = None):
from diffusers.utils import remote_utils
from modules import devices, shared, errors, modelloader
if not shared.opts.remote_vae_encode:
+1 -1
View File
@@ -11,7 +11,7 @@ import installer
gr_height = 512
max_units = shared.opts.control_max_units
units: list[unit.Unit] = [] # main state variable
controls: list[gr.component] = [] # list of gr controls
controls: list[gr.components.Component] = [] # list of gr controls
debug = shared.log.trace if os.environ.get('SD_CONTROL_DEBUG', None) is not None else lambda *args, **kwargs: None
debug('Trace: CONTROL')
+10 -7
View File
@@ -1,5 +1,6 @@
import os
import gradio as gr
from typing import TYPE_CHECKING, cast
from modules import errors
from modules.ui_components import ToolButton
@@ -90,10 +91,11 @@ class UiLoadsave:
apply_field(x, 'value', check_dropdown, getattr(x, 'init_field', None))
def check_tab_id(tab_id):
tab_items = list(filter(lambda e: isinstance(e, gr.TabItem), x.children))
if TYPE_CHECKING:
assert isinstance(x, gr.Tabs)
tab_items = cast('list[gr.TabItem]', list(filter(lambda e: isinstance(e, gr.TabItem), x.children))) # Force static type checker to get correct type
if type(tab_id) == str:
tab_ids = [t.id for t in tab_items]
return tab_id in tab_ids
return tab_id in [t.id for t in tab_items]
elif type(tab_id) == int:
return 0 <= tab_id < len(tab_items)
else:
@@ -290,11 +292,12 @@ class UiLoadsave:
self.ui_defaults_review = gr.HTML("", elem_id="ui_defaults_review")
def setup_ui(self):
review = [self.ui_defaults_review] if self.ui_defaults_review is not None else None
if self.ui_defaults_view:
self.ui_defaults_view.click(fn=self.ui_view, inputs=list(self.component_mapping.values()), outputs=[self.ui_defaults_review])
self.ui_defaults_view.click(fn=self.ui_view, inputs=list(self.component_mapping.values()), outputs=review)
if self.ui_defaults_apply:
self.ui_defaults_apply.click(fn=self.ui_apply, inputs=list(self.component_mapping.values()), outputs=[self.ui_defaults_review])
self.ui_defaults_apply.click(fn=self.ui_apply, inputs=list(self.component_mapping.values()), outputs=review)
if self.ui_defaults_restore:
self.ui_defaults_restore.click(fn=self.ui_restore, inputs=[], outputs=[self.ui_defaults_review])
self.ui_defaults_restore.click(fn=self.ui_restore, inputs=[], outputs=review)
if self.ui_defaults_submenu:
self.ui_defaults_submenu.click(fn=self.ui_submenu_apply, _js='uiOpenSubmenus', inputs=[self.ui_defaults_review], outputs=[self.ui_defaults_review])
self.ui_defaults_submenu.click(fn=self.ui_submenu_apply, _js='uiOpenSubmenus', inputs=review, outputs=review)
+10 -5
View File
@@ -1,5 +1,6 @@
import os
import inspect
from typing import cast
import gradio as gr
from modules import errors, sd_models, sd_vae, extras, sd_samplers, ui_symbols, modelstats
from modules.ui_components import ToolButton
@@ -143,7 +144,7 @@ def create_ui():
model_table = gr.HTML(value='', elem_id="model_list_table")
model_checkhash_btn.click(fn=sd_models.update_model_hashes, inputs=[], outputs=[model_table])
model_list_btn.click(fn=lambda: create_models_table(sd_models.checkpoints_list.values()), inputs=[], outputs=[model_table])
model_list_btn.click(fn=lambda: create_models_table(list(sd_models.checkpoints_list.values())), inputs=[], outputs=[model_table])
with gr.Tab(label="Metadata", elem_id="models_metadata_tab"):
from modules.civitai.metadata_civitai import civit_search_metadata, civit_update_metadata
@@ -178,7 +179,7 @@ def create_ui():
custom_name = gr.Textbox(label="New model name")
with gr.Row():
merge_mode = gr.Dropdown(choices=merge_methods.__all__, value="weighted_sum", label="Interpolation Method")
merge_mode_docs = gr.HTML(value=getattr(merge_methods, "weighted_sum", "").__doc__.replace("\n", "<br>"))
merge_mode_docs = gr.HTML(value=merge_methods.weighted_sum.__doc__.strip().replace("\n", "<br>")) # pylint: disable=no-member # pyright: ignore[reportOptionalMemberAccess]
with gr.Row():
primary_model_name = gr.Dropdown(sd_model_choices(), label="Primary model", value="None")
create_refresh_button(primary_model_name, sd_models.list_models, lambda: {"choices": sd_model_choices()}, "checkpoint_A_refresh")
@@ -318,7 +319,11 @@ def create_ui():
return gr.Slider.update(value=None, visible=False)
def show_help(mode):
doc = getattr(merge_methods, mode).__doc__.replace("\n", "<br>")
try:
doc = getattr(merge_methods, mode).__doc__.strip().replace("\n", "<br>")
except AttributeError:
log.warning(f'Merge mode "{mode}" is missing documentation')
doc = "Error: Documentation missing"
return gr.update(value=doc, visible=True)
def show_unload(device):
@@ -358,8 +363,8 @@ def create_ui():
merge_mode.input(fn=tertiary, inputs=merge_mode, outputs=[tertiary_model_name, tertiary_refresh])
merge_mode.input(fn=beta_visibility, inputs=merge_mode, outputs=[beta, alpha_label, beta_label, beta_apply_preset, beta_preset, beta_base, beta_in_blocks, beta_mid_block, beta_out_blocks])
re_basin.change(fn=show_iters, inputs=re_basin, outputs=re_basin_iterations)
apply_preset.click(fn=load_presets, inputs=[alpha_preset, alpha_preset_lambda], outputs=[alpha_base, alpha_in_blocks, alpha_mid_block, alpha_out_blocks, tabs])
beta_apply_preset.click(fn=load_presets, inputs=[beta_preset, beta_preset_lambda], outputs=[beta_base, beta_in_blocks, beta_mid_block, beta_out_blocks, tabs])
apply_preset.click(fn=load_presets, inputs=[alpha_preset, alpha_preset_lambda], outputs=[alpha_base, alpha_in_blocks, alpha_mid_block, alpha_out_blocks, cast("gr.components.Component", tabs)]) # Casting because Tabs has an update method.
beta_apply_preset.click(fn=load_presets, inputs=[beta_preset, beta_preset_lambda], outputs=[beta_base, beta_in_blocks, beta_mid_block, beta_out_blocks, cast("gr.components.Component", tabs)]) # Casting because Tabs has an update method.
modelmerger_merge.click(
fn=wrap_gradio_gpu_call(modelmerger, extra_outputs=lambda: [gr.update() for _ in range(4)], name='Models'),