add qol to the SD3 text encoder patcher nodes

This commit is contained in:
cubiq
2024-10-22 19:45:32 +02:00
parent f320ada613
commit 5b75a0e84b
2 changed files with 160 additions and 1 deletions
+83 -1
View File
@@ -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
View File
@@ -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",
}