This commit is contained in:
Alex "mcmonkey" Goodwin
2023-02-12 23:51:50 -08:00
parent bacb5c3dd8
commit eb93219f8c
+8 -33
View File
@@ -1,21 +1,12 @@
from pathlib import Path
from modules import scripts
from scripts.dynamic_thresholding import VALID_MODES as scheduler_list
from scripts.dynamic_thresholding import VALID_MODES
def find_module(name_list):
if isinstance(name_list, str):
name_list = [s.strip() for s in name_list.split(",")]
def make_axis_options():
xyz_grid = [x for x in scripts.scripts_data if x.script_class.__module__ == "xyz_grid.py"][0].module
for data in scripts.scripts_data:
if Path(data.path).name in name_list:
return data.module
return None
def make_axis_options(xyz_grid):
AxisOption = xyz_grid.AxisOption
apply_field = xyz_grid.apply_field
@@ -31,7 +22,7 @@ def make_axis_options(xyz_grid):
def apply_scheduler(field):
def core(p, x, xs):
if x not in scheduler_list:
if x not in VALID_MODES:
raise RuntimeError(f"Unknown Scheduler: {x}")
setattr(p, field, x)
@@ -41,30 +32,14 @@ def make_axis_options(xyz_grid):
extra_axis_options = [
AxisOption("DT Mimic Scale", float, apply_mimic_scale()),
AxisOption("DT Threshold", float, apply_field("dynthres_threshold_percentile")),
AxisOption("DT Mimic Scheduler", str, apply_scheduler("dynthres_mimic_mode"), choices=lambda: scheduler_list),
AxisOption("DT Mimic Scheduler", str, apply_scheduler("dynthres_mimic_mode"), choices=lambda: VALID_MODES),
AxisOption("DT Mimic min", float, apply_field("dynthres_mimic_scale_min")),
AxisOption("DT CFG Scheduler", str, apply_scheduler("dynthres_cfg_mode"), choices=lambda: scheduler_list),
AxisOption("DT CFG Scheduler", str, apply_scheduler("dynthres_cfg_mode"), choices=lambda: VALID_MODES),
AxisOption("DT CFG min", float, apply_field("dynthres_cfg_scale_min")),
AxisOption("DT Power val", float, apply_field("dynthres_power_val")),
AxisOption("DT Experiment", int, apply_field("dynthres_experiment_mode")),
]
return extra_axis_options
xyz_grid.axis_options.extend(extra_axis_options)
def add_axis_options(axis_options, extra_axis_options, target_label=""):
if target_label:
for i, axis_option in enumerate(axis_options):
if axis_option.label == target_label:
axis_options[i+1:i+1] = extra_axis_options
return
axis_options.extend(extra_axis_options)
return
xyz_grid = find_module("xyz_grid.py, xy_grid.py")
if xyz_grid:
extra_axis_options = make_axis_options(xyz_grid)
add_axis_options(xyz_grid.axis_options, extra_axis_options, "CFG Scale")
make_axis_options()