From 569d6858773fc0cdbd8f34bd3e00dd0adf243d37 Mon Sep 17 00:00:00 2001 From: blepping Date: Sat, 3 Feb 2024 12:25:50 -0700 Subject: [PATCH] Add sonar samplers to the built in samplers list --- __init__.py | 2 ++ py/sonar.py | 40 ++++++++++++++++++++++++++++++++++++---- 2 files changed, 38 insertions(+), 4 deletions(-) diff --git a/__init__.py b/__init__.py index 3365805..5443f40 100644 --- a/__init__.py +++ b/__init__.py @@ -1,5 +1,7 @@ from .py import sonar +sonar.add_samplers() + NODE_CLASS_MAPPINGS = { "SamplerSonarEuler": sonar.SamplerNodeSonarEuler, "SamplerSonarEulerA": sonar.SamplerNodeSonarEulerAncestral, diff --git a/py/sonar.py b/py/sonar.py index ba118c6..f6f6b7f 100644 --- a/py/sonar.py +++ b/py/sonar.py @@ -22,13 +22,13 @@ class HistoryType(Enum): class SonarBase: def __init__( self, - history_type: HistoryType = HistoryType.ZERO, + history_type: HistoryType | None = None, momentum: float = 0.95, momentum_hist: float = 0.75, direction: float = 1.0, ) -> None: self.history_d = None - self.history_type = history_type + self.history_type = HistoryType.ZERO if history_type is None else history_type self.momentum = momentum self.momentum_hist = momentum_hist self.direction = direction @@ -43,6 +43,8 @@ class SonarBase: self.history_d = x elif self.history_type == HistoryType.RAND: self.history_d = torch.randn_like(x) + else: + raise ValueError("Sonar sampler: bad history type") def momentum_step(self, x: Tensor, d: Tensor, dt: Tensor): hd = self.history_d @@ -142,7 +144,7 @@ class SonarEuler(SonarSampler): disable=None, momentum=0.95, momentum_hist=0.75, - momentum_init="ZERO", + momentum_init=HistoryType.ZERO, direction=1.0, s_churn=0.0, s_tmin=0.0, @@ -241,7 +243,7 @@ class SonarEulerAncestral(SonarSampler): disable=None, momentum=0.95, momentum_hist=0.75, - momentum_init="ZERO", + momentum_init=HistoryType.ZERO, noise_type="gaussian", direction=1.0, eta=1.0, @@ -417,3 +419,33 @@ class SamplerNodeSonarEulerAncestral(SamplerNodeSonarEuler): }, ), ) + + +def add_samplers(): + import importlib + + from comfy.samplers import KSampler, k_diffusion_sampling + + extra_samplers = { + "sonar_euler": SonarEuler.sampler, + "sonar_euler_ancestral": SonarEulerAncestral.sampler, + } + added = 0 + for ( + name, + sampler, + ) in extra_samplers.items(): + if name in KSampler.SAMPLERS: + continue + try: + KSampler.SAMPLERS.append(name) + setattr( + k_diffusion_sampling, + f"sample_{name}", + sampler, + ) + added += 1 + except ValueError as exc: + print(f"Sonar: Failed to add {name} to built in samplers list: {exc}") + if added > 0: + importlib.reload(k_diffusion_sampling)