better scaling

This commit is contained in:
Extraltodeus
2024-06-19 07:26:52 +02:00
committed by GitHub
parent 72bc6dea74
commit 66cd873ad3
+12 -4
View File
@@ -1,11 +1,14 @@
import torch
def get_scale(scale,sigmas,sigmin,sigmax):
current = (1 + ((sigmas - sigmax) * (0 - 1)) / (sigmin - sigmax)) ** .5 # forgetting the global parenthesis made me generate a image looking like a negative
return scale * current + (1 - current)
class uncondZeroNode:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"model": ("MODEL",),
"scale": ("FLOAT", {"default": 1, "min": 0.0, "max": 10.0, "step": 0.01, "round": 0.01}),
"method":(["uncond_zero","rescale_cfg"],)
"method":(["uncond_zero","rescale_cfg"],),
}}
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
@@ -13,12 +16,15 @@ class uncondZeroNode:
CATEGORY = "model_patches"
def patch(self, model, scale, method):
sigmin = model.model.model_sampling.sigma(model.model.model_sampling.timestep(model.model.model_sampling.sigma_min)).item()
sigmax = model.model.model_sampling.sigma(model.model.model_sampling.timestep(model.model.model_sampling.sigma_max)).item()
def uncond_zero(args):
cond = args["cond_denoised"]
x_orig = args["input"]
x_orig -= x_orig.mean()
cond -= cond.mean()
return x_orig - cond / cond.std() * scale
current_scale = get_scale(scale,args['sigma'][0].item(), sigmin, sigmax)
return x_orig - cond / cond.std() ** .5 * current_scale
new_scale = 1 / (model.model.latent_format.scale_factor * 8)
@@ -29,7 +35,9 @@ class uncondZeroNode:
cond = args["cond_denoised"]
cond -= cond.mean()
cond = x_orig - cond / cond.std() * scale
current_scale = get_scale(scale,args['sigma'][0].item(), sigmin, sigmax)
cond = x_orig - cond / cond.std() ** .5 * current_scale
sigma = args["sigma"]
sigma = sigma.view(sigma.shape[:1] + (1,) * (cond.ndim - 1))
@@ -50,4 +58,4 @@ class uncondZeroNode:
m = model.clone()
m.set_model_sampler_cfg_function({"uncond_zero":uncond_zero,"rescale_cfg":rescale_cfg}[method])
return (m, )
return (m, )