add qol to the SD3 text encoder patcher nodes
This commit is contained in:
@@ -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;
|
||||
}
|
||||
});
|
||||
});
|
||||
};
|
||||
}
|
||||
},
|
||||
});
|
||||
+77
@@ -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",
|
||||
}
|
||||
Reference in New Issue
Block a user