From 5b75a0e84ba92f944d801b2c9dfbf9940abdeba3 Mon Sep 17 00:00:00 2001 From: cubiq Date: Tue, 22 Oct 2024 19:45:32 +0200 Subject: [PATCH] add qol to the SD3 text encoder patcher nodes --- js/FluxAttentionSeeker.js | 84 ++++++++++++++++++++++++++++++++++++++- sampling.py | 77 +++++++++++++++++++++++++++++++++++ 2 files changed, 160 insertions(+), 1 deletion(-) diff --git a/js/FluxAttentionSeeker.js b/js/FluxAttentionSeeker.js index 4a7733f..b31989c 100644 --- a/js/FluxAttentionSeeker.js +++ b/js/FluxAttentionSeeker.js @@ -11,7 +11,6 @@ app.registerExtension({ const onCreated = nodeType.prototype.onNodeCreated; nodeType.prototype.onNodeCreated = function () { - console.log(LGraphCanvas); this.addWidget("button", "RESET ALL", null, () => { this.widgets.forEach(w => { if (w.type === "slider") { @@ -48,4 +47,87 @@ app.registerExtension({ }; } }, +}); + +app.registerExtension({ + name: "essentials.SD3AttentionSeekerLG", + async beforeRegisterNodeDef(nodeType, nodeData, app) { + if (!nodeData?.category?.startsWith("essentials")) { + return; + } + + if (nodeData.name === "SD3AttentionSeekerLG+") { + const onCreated = nodeType.prototype.onNodeCreated; + + nodeType.prototype.onNodeCreated = function () { + this.addWidget("button", "RESET L", null, () => { + this.widgets.forEach(w => { + if (w.type === "slider" && w.name.startsWith('clip_l')) { + w.value = 1.0; + } + }); + }); + this.addWidget("button", "RESET G", null, () => { + this.widgets.forEach(w => { + if (w.type === "slider" && w.name.startsWith('clip_g')) { + w.value = 1.0; + } + }); + }); + + this.addWidget("button", "REPEAT FIRST", null, () => { + var clip_l_value = undefined; + var clip_g_value = undefined; + this.widgets.forEach(w => { + if (w.name.startsWith('clip_l')) { + if (clip_l_value === undefined) { + clip_l_value = w.value; + } + w.value = clip_l_value; + } else if (w.name.startsWith('clip_g')) { + if (clip_g_value === undefined) { + clip_g_value = w.value; + } + w.value = clip_g_value; + } + }); + }); + }; + } + }, +}); + +app.registerExtension({ + name: "essentials.SD3AttentionSeekerT5", + async beforeRegisterNodeDef(nodeType, nodeData, app) { + if (!nodeData?.category?.startsWith("essentials")) { + return; + } + + if (nodeData.name === "SD3AttentionSeekerT5+") { + const onCreated = nodeType.prototype.onNodeCreated; + + nodeType.prototype.onNodeCreated = function () { + this.addWidget("button", "RESET ALL", null, () => { + this.widgets.forEach(w => { + if (w.type === "slider") { + w.value = 1.0; + } + }); + }); + + this.addWidget("button", "REPEAT FIRST", null, () => { + var t5_value = undefined; + this.widgets.forEach(w => { + if (w.name.startsWith('t5')) { + if (t5_value === undefined) { + t5_value = w.value; + } + w.value = t5_value; + } + }); + }); + }; + } + }, }); \ No newline at end of file diff --git a/sampling.py b/sampling.py index ef31534..f4244ea 100644 --- a/sampling.py +++ b/sampling.py @@ -707,6 +707,81 @@ class GuidanceTimestepping: m.set_model_sampler_cfg_function(apply_apg) return (m,) +class ModelSamplingDiscreteFlowCustom(torch.nn.Module): + def __init__(self, model_config=None): + super().__init__() + if model_config is not None: + sampling_settings = model_config.sampling_settings + else: + sampling_settings = {} + + self.set_parameters(shift=sampling_settings.get("shift", 1.0), multiplier=sampling_settings.get("multiplier", 1000)) + + def set_parameters(self, shift=1.0, timesteps=1000, multiplier=1000, cut_off=1.0, shift_multiplier=0): + self.shift = shift + self.multiplier = multiplier + self.cut_off = cut_off + self.shift_multiplier = shift_multiplier + ts = self.sigma((torch.arange(1, timesteps + 1, 1) / timesteps) * multiplier) + self.register_buffer('sigmas', ts) + + @property + def sigma_min(self): + return self.sigmas[0] + + @property + def sigma_max(self): + return self.sigmas[-1] + + def timestep(self, sigma): + return sigma * self.multiplier + + def sigma(self, timestep): + shift = self.shift + if timestep.dim() == 0: + t = timestep.cpu().item() / self.multiplier + if t <= self.cut_off: + shift = shift * self.shift_multiplier + + return comfy.model_sampling.time_snr_shift(shift, timestep / self.multiplier) + + def percent_to_sigma(self, percent): + if percent <= 0.0: + return 1.0 + if percent >= 1.0: + return 0.0 + return 1.0 - percent + +class ModelSamplingSD3Advanced: + @classmethod + def INPUT_TYPES(s): + return {"required": { "model": ("MODEL",), + "shift": ("FLOAT", {"default": 3.0, "min": 0.0, "max": 100.0, "step":0.01}), + "cut_off": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step":0.05}), + "shift_multiplier": ("FLOAT", {"default": 2, "min": 0, "max": 10, "step":0.05}), + }} + + RETURN_TYPES = ("MODEL",) + FUNCTION = "execute" + + CATEGORY = "essentials/sampling" + + def execute(self, model, shift, multiplier=1000, cut_off=1.0, shift_multiplier=0): + m = model.clone() + + + sampling_base = ModelSamplingDiscreteFlowCustom + sampling_type = comfy.model_sampling.CONST + + class ModelSamplingAdvanced(sampling_base, sampling_type): + pass + + model_sampling = ModelSamplingAdvanced(model.model.model_config) + model_sampling.set_parameters(shift=shift, multiplier=multiplier, cut_off=cut_off, shift_multiplier=shift_multiplier) + m.add_object_patch("model_sampling", model_sampling) + + return (m, ) + SAMPLING_CLASS_MAPPINGS = { "KSamplerVariationsStochastic+": KSamplerVariationsStochastic, "KSamplerVariationsWithNoise+": KSamplerVariationsWithNoise, @@ -718,6 +793,7 @@ SAMPLING_CLASS_MAPPINGS = { "SamplerSelectHelper+": SamplerSelectHelper, "SchedulerSelectHelper+": SchedulerSelectHelper, "LorasForFluxParams+": LorasForFluxParams, + "ModelSamplingSD3Advanced+": ModelSamplingSD3Advanced, } SAMPLING_NAME_MAPPINGS = { @@ -731,4 +807,5 @@ SAMPLING_NAME_MAPPINGS = { "SamplerSelectHelper+": "🔧 Sampler Select Helper", "SchedulerSelectHelper+": "🔧 Scheduler Select Helper", "LorasForFluxParams+": "🔧 LoRA for Flux Parameters", + "ModelSamplingSD3Advanced+": "🔧 Model Sampling SD3 Advanced", } \ No newline at end of file