307 lines
7.9 KiB
Python
307 lines
7.9 KiB
Python
from __future__ import annotations
|
|
|
|
import torch
|
|
from comfy import samplers
|
|
|
|
from . import noise
|
|
from .sonar import (
|
|
GuidanceConfig,
|
|
GuidanceType,
|
|
HistoryType,
|
|
SonarConfig,
|
|
SonarDPMPPSDE,
|
|
SonarEuler,
|
|
SonarEulerAncestral,
|
|
)
|
|
|
|
|
|
class NoisyLatentLikeNode:
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {
|
|
"required": {
|
|
"noise_type": (
|
|
tuple(
|
|
t.name.lower()
|
|
for t in noise.NoiseType
|
|
if t is not noise.NoiseType.BROWNIAN
|
|
),
|
|
),
|
|
"seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}),
|
|
"latent": ("LATENT",),
|
|
},
|
|
}
|
|
|
|
RETURN_TYPES = ("LATENT",)
|
|
CATEGORY = "latent/noise"
|
|
|
|
FUNCTION = "go"
|
|
|
|
def go(
|
|
self,
|
|
noise_type,
|
|
seed,
|
|
latent,
|
|
):
|
|
ns = noise.get_noise_sampler(
|
|
noise.NoiseType[noise_type.upper()],
|
|
latent["samples"],
|
|
None,
|
|
None,
|
|
seed=None,
|
|
use_cpu=True,
|
|
)
|
|
randst = torch.random.get_rng_state()
|
|
try:
|
|
torch.random.manual_seed(seed)
|
|
result = ns(None, None)
|
|
finally:
|
|
torch.random.set_rng_state(randst)
|
|
return ({"samples": result},)
|
|
|
|
|
|
class GuidanceConfigNode:
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {
|
|
"required": {
|
|
"factor": (
|
|
"FLOAT",
|
|
{
|
|
"default": 0.01,
|
|
"min": -2.0,
|
|
"max": 2.0,
|
|
"step": 0.001,
|
|
"round": False,
|
|
},
|
|
),
|
|
"guidance_type": (tuple(t.name.lower() for t in GuidanceType),),
|
|
"start_step": ("INT", {"default": 1, "min": 1}),
|
|
"end_step": ("INT", {"default": 9999, "min": 1}),
|
|
"latent": ("LATENT",),
|
|
},
|
|
}
|
|
|
|
RETURN_TYPES = ("SONAR_GUIDANCE_CFG",)
|
|
CATEGORY = "sampling/custom_sampling/samplers"
|
|
|
|
FUNCTION = "make_guidance_cfg"
|
|
|
|
def make_guidance_cfg(
|
|
self,
|
|
guidance_type,
|
|
factor,
|
|
start_step,
|
|
end_step,
|
|
latent,
|
|
):
|
|
return (
|
|
GuidanceConfig(
|
|
guidance_type=GuidanceType[guidance_type.upper()],
|
|
factor=factor,
|
|
start_step=start_step,
|
|
end_step=end_step,
|
|
latent=latent.get("samples"),
|
|
),
|
|
)
|
|
|
|
|
|
class SamplerNodeSonarBase:
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {
|
|
"required": {
|
|
"momentum": (
|
|
"FLOAT",
|
|
{
|
|
"default": 0.95,
|
|
"min": -0.5,
|
|
"max": 2.5,
|
|
"step": 0.01,
|
|
"round": False,
|
|
},
|
|
),
|
|
"momentum_hist": (
|
|
"FLOAT",
|
|
{
|
|
"default": 0.75,
|
|
"min": -1.5,
|
|
"max": 1.5,
|
|
"step": 0.01,
|
|
"round": False,
|
|
},
|
|
),
|
|
"momentum_init": (tuple(t.name for t in HistoryType),),
|
|
"direction": (
|
|
"FLOAT",
|
|
{
|
|
"default": 1.0,
|
|
"min": -30.0,
|
|
"max": 15.0,
|
|
"step": 0.01,
|
|
"round": False,
|
|
},
|
|
),
|
|
},
|
|
"optional": {"guidance_cfg_opt": ("SONAR_GUIDANCE_CFG",)},
|
|
}
|
|
|
|
RETURN_TYPES = ("SAMPLER",)
|
|
CATEGORY = "sampling/custom_sampling/samplers"
|
|
|
|
|
|
class SamplerNodeSonarEuler(SamplerNodeSonarBase):
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
result = super().INPUT_TYPES()
|
|
result["required"].update(
|
|
{
|
|
"s_noise": (
|
|
"FLOAT",
|
|
{
|
|
"default": 1.0,
|
|
"min": 0.0,
|
|
"max": 100.0,
|
|
"step": 0.01,
|
|
"round": False,
|
|
},
|
|
),
|
|
},
|
|
)
|
|
return result
|
|
|
|
RETURN_TYPES = ("SAMPLER",)
|
|
CATEGORY = "sampling/custom_sampling/samplers"
|
|
|
|
FUNCTION = "get_sampler"
|
|
|
|
def get_sampler(
|
|
self,
|
|
momentum,
|
|
momentum_hist,
|
|
momentum_init,
|
|
direction,
|
|
s_noise,
|
|
guidance_cfg_opt=None,
|
|
):
|
|
cfg = SonarConfig(
|
|
momentum=momentum,
|
|
init=HistoryType[momentum_init.upper()],
|
|
momentum_hist=momentum_hist,
|
|
direction=direction,
|
|
guidance=guidance_cfg_opt,
|
|
)
|
|
return (
|
|
samplers.KSAMPLER(
|
|
SonarEuler.sampler,
|
|
{
|
|
"s_noise": s_noise,
|
|
"sonar_config": cfg,
|
|
},
|
|
),
|
|
)
|
|
|
|
|
|
class SamplerNodeSonarEulerAncestral(SamplerNodeSonarEuler):
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
result = super().INPUT_TYPES()
|
|
result["required"].update(
|
|
{
|
|
"eta": (
|
|
"FLOAT",
|
|
{
|
|
"default": 1.0,
|
|
"min": 0.0,
|
|
"max": 100.0,
|
|
"step": 0.01,
|
|
"round": False,
|
|
},
|
|
),
|
|
"noise_type": (tuple(t.name.lower() for t in noise.NoiseType),),
|
|
},
|
|
)
|
|
return result
|
|
|
|
def get_sampler(
|
|
self,
|
|
momentum,
|
|
momentum_hist,
|
|
momentum_init,
|
|
direction,
|
|
noise_type,
|
|
eta,
|
|
s_noise,
|
|
guidance_cfg_opt=None,
|
|
):
|
|
cfg = SonarConfig(
|
|
momentum=momentum,
|
|
init=HistoryType[momentum_init.upper()],
|
|
momentum_hist=momentum_hist,
|
|
direction=direction,
|
|
noise_type=noise.NoiseType[noise_type.upper()],
|
|
guidance=guidance_cfg_opt,
|
|
)
|
|
return (
|
|
samplers.KSAMPLER(
|
|
SonarEulerAncestral.sampler,
|
|
{
|
|
"sonar_config": cfg,
|
|
"eta": eta,
|
|
"s_noise": s_noise,
|
|
},
|
|
),
|
|
)
|
|
|
|
|
|
class SamplerNodeSonarDPMPPSDE(SamplerNodeSonarEuler):
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
result = super().INPUT_TYPES()
|
|
result["required"].update(
|
|
{
|
|
"eta": (
|
|
"FLOAT",
|
|
{
|
|
"default": 1.0,
|
|
"min": 0.0,
|
|
"max": 100.0,
|
|
"step": 0.01,
|
|
"round": False,
|
|
},
|
|
),
|
|
"noise_type": (tuple(t.name.lower() for t in noise.NoiseType),),
|
|
},
|
|
)
|
|
return result
|
|
|
|
def get_sampler(
|
|
self,
|
|
momentum,
|
|
momentum_hist,
|
|
momentum_init,
|
|
direction,
|
|
noise_type,
|
|
eta,
|
|
s_noise,
|
|
guidance_cfg_opt=None,
|
|
):
|
|
cfg = SonarConfig(
|
|
momentum=momentum,
|
|
init=HistoryType[momentum_init.upper()],
|
|
momentum_hist=momentum_hist,
|
|
direction=direction,
|
|
noise_type=noise.NoiseType[noise_type.upper()],
|
|
guidance=guidance_cfg_opt,
|
|
)
|
|
return (
|
|
samplers.KSAMPLER(
|
|
SonarDPMPPSDE.sampler,
|
|
{
|
|
"sonar_config": cfg,
|
|
"eta": eta,
|
|
"s_noise": s_noise,
|
|
},
|
|
),
|
|
)
|