Compare commits
47
Commits
ddim_removed
...
main
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
41ce3ce693 | ||
|
|
25d8be7160 | ||
|
|
5b06f9623f | ||
|
|
985fc1f3e8 | ||
|
|
e7c6ffe3c0 | ||
|
|
b00812334b | ||
|
|
9d65ad1b68 | ||
|
|
d17907e09d | ||
|
|
d12cb9e96b | ||
|
|
725d3aeb58 | ||
|
|
83470d59e0 | ||
|
|
7ef2f56798 | ||
|
|
a36442dd4f | ||
|
|
1b65c60e0c | ||
|
|
f5c7a877e5 | ||
|
|
471b5972c9 | ||
|
|
20dc195ca9 | ||
|
|
bda63be15c | ||
|
|
e2dcfd4091 | ||
|
|
2669b01670 | ||
|
|
3752a2daf9 | ||
|
|
ea79890408 | ||
|
|
a73142188a | ||
|
|
bb292f8ec1 | ||
|
|
bc784ec621 | ||
|
|
ae32ce995a | ||
|
|
eca26065e4 | ||
|
|
9c65214933 | ||
|
|
f890b8bf62 | ||
|
|
17193ecdbf | ||
|
|
e2e25e7a28 | ||
|
|
bbbddbd7cb | ||
|
|
67d4b62235 | ||
|
|
12c0ca1044 | ||
|
|
b172908ac7 | ||
|
|
33ba61bb78 | ||
|
|
22f4ee04e5 | ||
|
|
b5b8999f5d | ||
|
|
39c5f85d80 | ||
|
|
b4d2fe9661 | ||
|
|
b5ef71262f | ||
|
|
eebbd9cf21 | ||
|
|
1fb1b03261 | ||
|
|
20efcba391 | ||
|
|
6ddbaf02f4 | ||
|
|
24c70dce64 | ||
|
|
f30cb0ef9c |
@@ -1,10 +1,12 @@
|
||||
# ComfyUI_restart_sampling
|
||||
Unofficial [ComfyUI](https://github.com/comfyanonymous/ComfyUI) nodes for restart sampling based on the paper "Restart Sampling for Improving Generative Processes"
|
||||
Unofficial [ComfyUI](https://github.com/comfyanonymous/ComfyUI) nodes for restart sampling based on the paper "Restart Sampling for Improving Generative Processes"
|
||||
|
||||
Paper: https://arxiv.org/abs/2306.14878
|
||||
|
||||
Repo: https://github.com/Newbeeer/diffusion_restart_sampling
|
||||
|
||||
This has been tested for ComfyUI for the following commit: [72508a8](https://github.com/comfyanonymous/ComfyUI/commit/72508a8d19121e2814ea4dfbce8a5311f37dcd61)
|
||||
|
||||
## Installation
|
||||
|
||||
Enter the following command from the commandline starting in ComfyUI/custom_nodes/
|
||||
@@ -14,26 +16,76 @@ git clone https://github.com/ssitu/ComfyUI_restart_sampling
|
||||
|
||||
## Usage
|
||||
|
||||
Nodes can be found in the node menu under `sampling`:
|
||||
The Restart sampler nodes can be found in the node menu under `sampling`.
|
||||
|
||||
If you set the environment variable `COMFYUI_VERBOSE_RESTART_SAMPLING` to `1`, restart sampling will dump
|
||||
information about the steps it's going to run to the console.
|
||||
|
||||
### Nodes
|
||||
|
||||
|Node|Image|Description|
|
||||
| --- | --- | --- |
|
||||
| KSampler With Restarts |  | Has all the inputs of a KSampler, but with an added string widget for configuring the Restart segments and a widget for the scheduler for the Restart segments. Not all samplers and schedulers from KSampler are currently supported. Restart sampling is done with ODE samplers and are not supposed to be used with SDE samplers. <br>The format for `segments` is a sequence of comma separated arrays of ${[N_{\textrm{Restart}}, K, t_{\textrm{min}}, t_{\textrm{max}}]}$. For example, [4, 1, 19.35, 40.79], [4, 1, 1.09, 1.92], [4, 5, 0.59, 1.09], [4, 5, 0.30, 0.59], [6, 6, 0.06, 0.30] would be a valid sequence. Segments may overwrite each other if their $t_{\textrm{min}}$ parameters are too close to each other. Each segment will add $(N_{\textrm{Restart}} - 1) \cdot K$ steps to the sampling process. For more information on the Restart parameters, refer to the paper. <br>The `restart_scheduler` is used as the scheduler for the denoising process during restart segments. The researchers used the Karras scheduler in their experiments, but use the same scheduler as the sampler schedule in their implementation. |
|
||||
| KSampler With Restarts |  | Has all the inputs of a KSampler, but with an added string widget for configuring the Restart segments and a widget for the scheduler for the Restart segments. Not all samplers and schedulers from KSampler are currently supported. Restart sampling is done with ODE samplers and are not supposed to be used with SDE samplers. <br>See the [Segments](#segments) section below for information how to define segments. For more information on the Restart parameters, refer to the paper. <br>The `restart_scheduler` is used as the scheduler for the denoising process during restart segments. The researchers used the Karras scheduler in their experiments, but use the same scheduler as the sampler schedule in their implementation. |
|
||||
| KSampler With Restarts (Simple) | | Instead of having a restart segment scheduler, segments will use the same scheduler as the KSampler scheduler. |
|
||||
| KSampler With Restarts (Advanced) | | Has all the inputs for an Advanced KSampler with all the inputs for restart sampling. It should be noted that there is a possibility for invalid segments when using it to end the denoising process early or starting it late (e.g. 20 steps, start at step 0, end at step 10) and invalid segments will be ignored. An invalid segment means that the closest $t_{\textrm{min}}$ in the noise schedule is higher than the segment's $t_{\textrm{max}}$, so the segment would have restarted the denoising process at $t_{\textrm{max}}$ then try to go to a higher noise level (when it should've gone to a lower noise level near $t_{\textrm{min}}$) which will destroy the sample. |
|
||||
| KSampler With Restarts (Custom) | | Essentially the same as `KSampler With Restarts (Advanced)` but it takes a `SAMPLER` input like the built in `SamplerCustom` node. Note that it is possible to input samplers that don't work properly or are incompatible with Restart sampling like SDE and UniPC samplers.|
|
||||
| `RestartScheduler` | | For use with custom sampling: This node will output sigmas like other scheduler nodes with restart segments inserted. Must be used with `RestartSampler`. Like stand alone samplers, the node takes parameters for restart segments and schedules. You may also optionally connect sigmas to it, in which case it will use the supplied sigmas for the main schedule. **Note**: When sigmas are connected, the `steps` and `scheduler` parameters have no effect. Setting `denoise` also can't adjust the steps: it can only shorten the sigmas you pass to the node. |
|
||||
| `RestartSampler` | | For use with custom sampling: Should be used in conjunction with `RestartScheduler` and takes a `SAMPLER` input. This node arranges for the restart noise to be injected at the appropriate points and delegates to the supplied sampler for actual sampling. |
|
||||
|
||||
---
|
||||
### Segments
|
||||
|
||||
## Comparisons
|
||||
These images can be dragged into ComfyUI to load their workflows.
|
||||
Each image is done using the Stable Diffusion v1.5 checkpoint with 18 steps using the Heun sampler and a Karras schedule. The images in the right column use a restart segment of $[N_{\textrm{Restart}}=3, K=2, t_{\textrm{min}}=0.06, t_{\textrm{max}}=0.30]$, which adds 4 steps.
|
||||
| Without | With |
|
||||
| --- | --- |
|
||||
|  |  |
|
||||
|  |  |
|
||||
|  |  |
|
||||
The format for `segments` is a sequence of comma separated arrays of ${[N_{\textrm{Restart}}, K, t_{\textrm{min}}, t_{\textrm{max}}]}$. For example, `[4, 1, 19.35, 40.79], [4, 1, 1.09, 1.92], [4, 5, 0.59, 1.09], [4, 5, 0.30, 0.59], [6, 6, 0.06, 0.30]` would be a valid sequence. Segments may overwrite each other if their $t_{\textrm{min}}$ parameters are too close to each other. Each segment will add $(N_{\textrm{Restart}} - 1) \cdot K$ steps to the sampling process.
|
||||
|
||||
Image slider links:
|
||||
- https://imgsli.com/MTg5NzI4
|
||||
- https://imgsli.com/MTg5NzI5
|
||||
- https://imgsli.com/MTg5NzI3
|
||||
Both $t_{\textrm{min}}$ and $t_{\textrm{max}}$ within a segment definition may be specified in any of the following three ways:
|
||||
|
||||
* A positive numeric value (i.e. `1.2`) — this will be interpreted as a sigma value.
|
||||
* A negative numeric value between `-0` and `-1000` — this will be interpeted as a (positive) timestep. Timesteps will be converted to integer values so if you need to specify timestep `0` you can do something like `-0.1`.
|
||||
* A quoted string percentage value followed by a percent sign (i.e. `"25%"`) — note that this refers to the percentage of sampling, not the percentage of steps that have elapsed.
|
||||
|
||||
You may freely mix the different formats. For example, `[2, 2, -500, "10%"], [3, 2, 5.3, -3]` would be a valid sequence. Note: Random numbers used for example only, not recommended.
|
||||
|
||||
**Special segment values**:
|
||||
|
||||
* Enter `default` by itself to use the default segment list.
|
||||
* Enter `a1111` by itself to emulate A1111 WebUI's segment calculation behavior. For full emulation, enabled chunked mode, set both schedulers to `karras` and the sampler to `heun`.
|
||||
* You may also enter `"default"` or `"a1111"` in place of a segment definition (note the quotes). This will insert preset segments at the point the quoted preset name appears. For example `[1,2,3,4], "default"` is the same as `[1,2,3,4], [3,2,0.06,0.30], [3,1,0.30,0.59]`.
|
||||
|
||||
### Chunked Mode
|
||||
|
||||
When chunked mode is enabled, the sampler is called with as many steps as possible up to the next segment. When disabled, the sampler
|
||||
is only called with a single step at a time. Some samplers such as SDE samplers, momentum samplers, second order samplers
|
||||
like dpmpp_2m use state from previous steps - when called step-by-step, this state is lost. Using chunked mode may make those
|
||||
samplers more accurate.
|
||||
|
||||
*Note*: Using SDE or momentum samplers with restart is likely not an improvement over normal sampling.
|
||||
|
||||
## Visual Example
|
||||
|
||||
Consider the default segments of `[3,2,0.06,0.30],[3,1,0.30,0.59]`.
|
||||
|
||||
1. $N_{\textrm{Restart}}=3, {K}=2, t_{\textrm{min}}=0.06, t_{\textrm{max}}=0.30$ — closer to the end of sampling, will run two restarts two times.
|
||||
2. $N_{\textrm{Restart}}=3, {K}=1, t_{\textrm{min}}=0.30, t_{\textrm{max}}=0.59$ — closer to the beginning of sampling, will run two restarts one time.
|
||||
|
||||
Running 20 steps with normal scheduling will look something like this:
|
||||
|
||||
```plaintext
|
||||
Step 1: sigma=10.7
|
||||
Step 2: sigma=8.08
|
||||
Step 3: sigma=6.2
|
||||
[... elided for brevity]
|
||||
Step 15: sigma=0.596
|
||||
Step 16: sigma=0.474
|
||||
Step 17: sigma=0.356 -- [3,1,0.30,0.59] matches here.
|
||||
K=1:
|
||||
restart 1, step 18
|
||||
restart 2, step 19
|
||||
Step 20: sigma=0.232
|
||||
Step 21: sigma=0.0292 -- [3,2,0.06,0.30] matches here.
|
||||
K=1:
|
||||
restart 1, step 22
|
||||
restart 2, step 23
|
||||
K=2:
|
||||
restart 1, step 24
|
||||
restart 2, step 25
|
||||
Step 26: sigma=0.0
|
||||
```
|
||||
|
||||
Binary file not shown.
|
Before Width: | Height: | Size: 1.1 MiB |
Binary file not shown.
|
Before Width: | Height: | Size: 468 KiB |
Binary file not shown.
|
Before Width: | Height: | Size: 518 KiB |
Binary file not shown.
|
Before Width: | Height: | Size: 1.1 MiB |
Binary file not shown.
|
Before Width: | Height: | Size: 470 KiB |
Binary file not shown.
|
Before Width: | Height: | Size: 523 KiB |
@@ -1,5 +1,20 @@
|
||||
import os
|
||||
|
||||
import comfy
|
||||
from .restart_sampling import restart_sampling, SCHEDULER_MAPPING
|
||||
|
||||
from .restart_sampling import (
|
||||
DEFAULT_SEGMENTS,
|
||||
NORMAL_SCHEDULER_MAPPING,
|
||||
RESTART_SCHEDULER_MAPPING,
|
||||
VERBOSE,
|
||||
RestartPlan,
|
||||
RestartSampler,
|
||||
restart_sampling,
|
||||
)
|
||||
|
||||
INCLUDE_SELFTEST = (
|
||||
os.environ.get("COMFYUI_RESTART_SAMPLING_SELFTEST", "").strip() == "1"
|
||||
)
|
||||
|
||||
|
||||
def get_supported_samplers():
|
||||
@@ -22,112 +37,400 @@ def get_supported_samplers():
|
||||
|
||||
|
||||
def get_supported_restart_schedulers():
|
||||
return list(SCHEDULER_MAPPING.keys())
|
||||
return tuple(RESTART_SCHEDULER_MAPPING.keys())
|
||||
|
||||
DEFAULT_SEGMENTS = "[3,2,0.06,0.30],[3,1,0.30,0.59]"
|
||||
|
||||
def get_supported_normal_schedulers():
|
||||
return tuple(NORMAL_SCHEDULER_MAPPING.keys())
|
||||
|
||||
|
||||
class KRestartSamplerSimple:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL", ),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"model": ("MODEL",),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}),
|
||||
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
|
||||
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0}),
|
||||
"sampler_name": (get_supported_samplers(), ),
|
||||
"scheduler": (comfy.samplers.KSampler.SCHEDULERS, ),
|
||||
"positive": ("CONDITIONING", ),
|
||||
"negative": ("CONDITIONING", ),
|
||||
"latent_image": ("LATENT", ),
|
||||
"denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
"segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}),
|
||||
}
|
||||
"sampler_name": (get_supported_samplers(),),
|
||||
"scheduler": (get_supported_restart_schedulers(),),
|
||||
"positive": ("CONDITIONING",),
|
||||
"negative": ("CONDITIONING",),
|
||||
"latent_image": ("LATENT",),
|
||||
"denoise": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01},
|
||||
),
|
||||
"segments": ("STRING", {"default": "default", "multiline": False}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "sampling"
|
||||
|
||||
def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments):
|
||||
return restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, segments, scheduler, denoise=denoise)
|
||||
def sample(
|
||||
self,
|
||||
model,
|
||||
seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler_name,
|
||||
scheduler,
|
||||
positive,
|
||||
negative,
|
||||
latent_image,
|
||||
denoise,
|
||||
segments,
|
||||
):
|
||||
return restart_sampling(
|
||||
model,
|
||||
seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler_name,
|
||||
scheduler,
|
||||
positive,
|
||||
negative,
|
||||
latent_image,
|
||||
segments,
|
||||
scheduler,
|
||||
denoise=denoise,
|
||||
)
|
||||
|
||||
|
||||
class KRestartSampler:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL", ),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
|
||||
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0}),
|
||||
"sampler_name": (get_supported_samplers(), ),
|
||||
"scheduler": (comfy.samplers.KSampler.SCHEDULERS, ),
|
||||
"positive": ("CONDITIONING", ),
|
||||
"negative": ("CONDITIONING", ),
|
||||
"latent_image": ("LATENT", ),
|
||||
"denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
"segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}),
|
||||
"restart_scheduler": (get_supported_restart_schedulers(), ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "sampling"
|
||||
|
||||
def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments, restart_scheduler):
|
||||
return restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, segments, restart_scheduler, denoise=denoise)
|
||||
|
||||
|
||||
class KRestartSamplerAdv:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"add_noise": (["enable", "disable"], ),
|
||||
"noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}),
|
||||
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
|
||||
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0}),
|
||||
"sampler_name": (get_supported_samplers(), ),
|
||||
"scheduler": (comfy.samplers.KSampler.SCHEDULERS, ),
|
||||
"positive": ("CONDITIONING", ),
|
||||
"negative": ("CONDITIONING", ),
|
||||
"latent_image": ("LATENT", ),
|
||||
"start_at_step": ("INT", {"default": 0, "min": 0, "max": 10000}),
|
||||
"end_at_step": ("INT", {"default": 10000, "min": 0, "max": 10000}),
|
||||
"return_with_leftover_noise": (["disable", "enable"], ),
|
||||
"segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}),
|
||||
"restart_scheduler": (get_supported_restart_schedulers(), ),
|
||||
}
|
||||
"sampler_name": (get_supported_samplers(),),
|
||||
"scheduler": (get_supported_normal_schedulers(),),
|
||||
"positive": ("CONDITIONING",),
|
||||
"negative": ("CONDITIONING",),
|
||||
"latent_image": ("LATENT",),
|
||||
"denoise": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01},
|
||||
),
|
||||
"segments": (
|
||||
"STRING",
|
||||
{"default": DEFAULT_SEGMENTS, "multiline": False},
|
||||
),
|
||||
"restart_scheduler": (get_supported_restart_schedulers(),),
|
||||
"chunked_mode": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "sampling"
|
||||
|
||||
def sample(self, model, add_noise, noise_seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, start_at_step, end_at_step, return_with_leftover_noise, segments, restart_scheduler):
|
||||
force_full_denoise = True
|
||||
if return_with_leftover_noise == "enable":
|
||||
force_full_denoise = False
|
||||
disable_noise = False
|
||||
if add_noise == "disable":
|
||||
disable_noise = True
|
||||
return restart_sampling(model, noise_seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, segments, restart_scheduler, disable_noise=disable_noise, start_step=start_at_step, last_step=end_at_step, force_full_denoise=force_full_denoise)
|
||||
def sample(
|
||||
self,
|
||||
model,
|
||||
seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler_name,
|
||||
scheduler,
|
||||
positive,
|
||||
negative,
|
||||
latent_image,
|
||||
denoise,
|
||||
segments,
|
||||
restart_scheduler,
|
||||
chunked_mode=True,
|
||||
):
|
||||
return restart_sampling(
|
||||
model,
|
||||
seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler_name,
|
||||
scheduler,
|
||||
positive,
|
||||
negative,
|
||||
latent_image,
|
||||
segments,
|
||||
restart_scheduler,
|
||||
denoise=denoise,
|
||||
chunked_mode=chunked_mode,
|
||||
)
|
||||
|
||||
|
||||
class KRestartSamplerAdv:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"add_noise": (["enable", "disable"],),
|
||||
"noise_seed": (
|
||||
"INT",
|
||||
{"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF},
|
||||
),
|
||||
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
|
||||
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0}),
|
||||
"sampler_name": (get_supported_samplers(),),
|
||||
"scheduler": (get_supported_normal_schedulers(),),
|
||||
"positive": ("CONDITIONING",),
|
||||
"negative": ("CONDITIONING",),
|
||||
"latent_image": ("LATENT",),
|
||||
"start_at_step": ("INT", {"default": 0, "min": 0, "max": 10000}),
|
||||
"end_at_step": ("INT", {"default": 10000, "min": 0, "max": 10000}),
|
||||
"return_with_leftover_noise": (["disable", "enable"],),
|
||||
"segments": (
|
||||
"STRING",
|
||||
{"default": DEFAULT_SEGMENTS, "multiline": False},
|
||||
),
|
||||
"restart_scheduler": (get_supported_restart_schedulers(),),
|
||||
"chunked_mode": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "sampling"
|
||||
|
||||
def sample(
|
||||
self,
|
||||
model,
|
||||
add_noise,
|
||||
noise_seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler_name,
|
||||
scheduler,
|
||||
positive,
|
||||
negative,
|
||||
latent_image,
|
||||
start_at_step,
|
||||
end_at_step,
|
||||
return_with_leftover_noise,
|
||||
segments,
|
||||
restart_scheduler,
|
||||
chunked_mode=True,
|
||||
):
|
||||
force_full_denoise = return_with_leftover_noise != "enable"
|
||||
disable_noise = add_noise == "disable"
|
||||
return restart_sampling(
|
||||
model,
|
||||
noise_seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler_name,
|
||||
scheduler,
|
||||
positive,
|
||||
negative,
|
||||
latent_image,
|
||||
segments,
|
||||
restart_scheduler,
|
||||
disable_noise=disable_noise,
|
||||
step_range=(start_at_step, end_at_step),
|
||||
force_full_denoise=force_full_denoise,
|
||||
chunked_mode=chunked_mode,
|
||||
)
|
||||
|
||||
|
||||
class KRestartSamplerCustom:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"add_noise": (["enable", "disable"],),
|
||||
"noise_seed": (
|
||||
"INT",
|
||||
{"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF},
|
||||
),
|
||||
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
|
||||
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0}),
|
||||
"sampler": ("SAMPLER",),
|
||||
"scheduler": (get_supported_normal_schedulers(),),
|
||||
"positive": ("CONDITIONING",),
|
||||
"negative": ("CONDITIONING",),
|
||||
"latent_image": ("LATENT",),
|
||||
"start_at_step": ("INT", {"default": 0, "min": 0, "max": 10000}),
|
||||
"end_at_step": ("INT", {"default": 10000, "min": 0, "max": 10000}),
|
||||
"return_with_leftover_noise": (["disable", "enable"],),
|
||||
"segments": (
|
||||
"STRING",
|
||||
{"default": DEFAULT_SEGMENTS, "multiline": False},
|
||||
),
|
||||
"restart_scheduler": (get_supported_restart_schedulers(),),
|
||||
"chunked_mode": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT", "LATENT")
|
||||
RETURN_NAMES = ("output", "denoised_output")
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "sampling"
|
||||
|
||||
def sample(
|
||||
self,
|
||||
model,
|
||||
add_noise,
|
||||
noise_seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler,
|
||||
scheduler,
|
||||
positive,
|
||||
negative,
|
||||
latent_image,
|
||||
start_at_step,
|
||||
end_at_step,
|
||||
return_with_leftover_noise,
|
||||
segments,
|
||||
restart_scheduler,
|
||||
chunked_mode=True,
|
||||
):
|
||||
force_full_denoise = return_with_leftover_noise != "enable"
|
||||
disable_noise = add_noise == "disable"
|
||||
return restart_sampling(
|
||||
model,
|
||||
noise_seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler,
|
||||
scheduler,
|
||||
positive,
|
||||
negative,
|
||||
latent_image,
|
||||
segments,
|
||||
restart_scheduler,
|
||||
disable_noise=disable_noise,
|
||||
step_range=(start_at_step, end_at_step),
|
||||
force_full_denoise=force_full_denoise,
|
||||
output_only=False,
|
||||
chunked_mode=chunked_mode,
|
||||
)
|
||||
|
||||
|
||||
class RestartSchedulerNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
|
||||
"scheduler": (get_supported_normal_schedulers(),),
|
||||
"segments": (
|
||||
"STRING",
|
||||
{"default": DEFAULT_SEGMENTS, "multiline": False},
|
||||
),
|
||||
"restart_scheduler": (get_supported_restart_schedulers(),),
|
||||
"denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0}),
|
||||
"start_at_step": ("INT", {"default": 0, "min": 0, "max": 10000}),
|
||||
"end_at_step": ("INT", {"default": 10000, "min": 0, "max": 10000}),
|
||||
},
|
||||
"optional": {
|
||||
"sigmas_opt": ("SIGMAS",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("SIGMAS",)
|
||||
FUNCTION = "go"
|
||||
CATEGORY = "sampling/custom_sampling/schedulers"
|
||||
|
||||
def go(
|
||||
self,
|
||||
model,
|
||||
steps,
|
||||
scheduler,
|
||||
segments,
|
||||
restart_scheduler,
|
||||
denoise,
|
||||
start_at_step=0,
|
||||
end_at_step=10000,
|
||||
sigmas_opt=None,
|
||||
):
|
||||
plan = RestartPlan(
|
||||
model,
|
||||
steps,
|
||||
scheduler,
|
||||
segments,
|
||||
restart_scheduler,
|
||||
denoise=denoise,
|
||||
step_range=(start_at_step, end_at_step),
|
||||
sigmas=sigmas_opt,
|
||||
)
|
||||
if VERBOSE:
|
||||
plan.explain(chunked=True)
|
||||
return (plan.sigmas(),)
|
||||
|
||||
|
||||
class RestartSamplerNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"sampler": ("SAMPLER",),
|
||||
"chunked_mode": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("SAMPLER",)
|
||||
FUNCTION = "go"
|
||||
CATEGORY = "sampling/custom_sampling/samplers"
|
||||
|
||||
def go(self, sampler, chunked_mode):
|
||||
restart_options = {
|
||||
"restart_chunked": chunked_mode,
|
||||
"restart_wrapped_sampler": sampler,
|
||||
}
|
||||
restart_sampler = comfy.samplers.KSAMPLER(
|
||||
RestartSampler.sampler_function,
|
||||
extra_options=sampler.extra_options | restart_options,
|
||||
inpaint_options=sampler.inpaint_options,
|
||||
)
|
||||
return (restart_sampler,)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"KRestartSamplerSimple": KRestartSamplerSimple,
|
||||
"KRestartSampler": KRestartSampler,
|
||||
"KRestartSamplerAdv": KRestartSamplerAdv,
|
||||
"KRestartSamplerCustom": KRestartSamplerCustom,
|
||||
"RestartScheduler": RestartSchedulerNode,
|
||||
"RestartSampler": RestartSamplerNode,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"KRestartSamplerSimple": "KSampler With Restarts (Simple)",
|
||||
"KRestartSampler": "KSampler With Restarts",
|
||||
"KRestartSamplerAdv": "KSampler With Restarts (Advanced)",
|
||||
"KRestartSamplerCustom": "KSampler With Restarts (Custom)",
|
||||
}
|
||||
|
||||
if INCLUDE_SELFTEST:
|
||||
|
||||
class RestartSelfTestNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"enabled": ("BOOLEAN", {"default": True}),
|
||||
"min_steps": ("INT", {"default": 2, "min": 0}),
|
||||
"max_steps": ("INT", {"default": 100, "min": 2}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("BOOLEAN",)
|
||||
FUNCTION = "go"
|
||||
CATEGORY = "sampling/custom_sampling/samplers"
|
||||
|
||||
def go(self, model, enabled=True, min_steps=2, max_steps=100):
|
||||
if enabled:
|
||||
RestartPlan.self_test(model, min_steps=min_steps, max_steps=max_steps)
|
||||
return (True,)
|
||||
|
||||
NODE_CLASS_MAPPINGS["RestartSelfTest"] = RestartSelfTestNode
|
||||
|
||||
+687
-86
@@ -1,33 +1,113 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
import os
|
||||
import warnings
|
||||
from collections import namedtuple
|
||||
from typing import NamedTuple
|
||||
|
||||
import comfy
|
||||
import latent_preview
|
||||
import torch
|
||||
from tqdm.auto import trange
|
||||
from nodes import common_ksampler
|
||||
from comfy.k_diffusion import sampling as k_diffusion_sampling
|
||||
from comfy import model_sampling
|
||||
from comfy.sample import prepare_noise, sample_custom
|
||||
from comfy.samplers import KSAMPLER, KSampler, sampler_object
|
||||
from comfy.utils import ProgressBar
|
||||
from .restart_schedulers import SCHEDULER_MAPPING
|
||||
from tqdm.auto import trange
|
||||
|
||||
from .restart_schedulers import NORMAL_SCHEDULER_MAPPING, RESTART_SCHEDULER_MAPPING
|
||||
|
||||
VERBOSE = os.environ.get("COMFYUI_VERBOSE_RESTART_SAMPLING", "").strip() == "1"
|
||||
|
||||
DEFAULT_SEGMENTS = "[3,2,0.06,0.30],[3,1,0.30,0.59]"
|
||||
|
||||
|
||||
def add_restart_segment(restart_segments, n_restart, k, t_min, t_max):
|
||||
if restart_segments is None:
|
||||
restart_segments = []
|
||||
restart_segments.append({'n': n_restart, 'k': k, 't_min': t_min, 't_max': t_max})
|
||||
restart_segments.append({"n": n_restart, "k": k, "t_min": t_min, "t_max": t_max})
|
||||
return restart_segments
|
||||
|
||||
|
||||
def prepare_restart_segments(restart_info):
|
||||
try:
|
||||
restart_arrays = ast.literal_eval(f"[{restart_info}]")
|
||||
except SyntaxError as e:
|
||||
print("Ill-formed restart segments")
|
||||
raise e
|
||||
def resolve_t_value(val, ms):
|
||||
if isinstance(val, (float, int)):
|
||||
if val >= 0.0:
|
||||
return val
|
||||
if val >= -1000:
|
||||
return ms.sigma(torch.FloatTensor([abs(int(val))], device="cpu")).item()
|
||||
if isinstance(val, str) and val.endswith("%"):
|
||||
try:
|
||||
val = float(val[:-1])
|
||||
if val >= 0 and val <= 100:
|
||||
return ms.percent_to_sigma(1.0 - val / 100.0)
|
||||
except ValueError:
|
||||
pass
|
||||
raise ValueError("bad t_min or t_max value")
|
||||
|
||||
|
||||
def prepare_restart_segments(restart_info, ms, sigmas):
|
||||
def get_a1111_segment():
|
||||
# Emulate A1111 WebUI's restart sampler behavior.
|
||||
steps = len(sigmas) - 1
|
||||
if steps < 20:
|
||||
# Less than 20 steps - no restarts.
|
||||
return []
|
||||
a1111_t_max = sigmas[int(torch.argmin(abs(sigmas - 2.0), dim=0))].item()
|
||||
if steps < 36:
|
||||
# Less than 36 steps - one restart with 9 steps.
|
||||
return [10, 1, 0.1, a1111_t_max]
|
||||
# Otherwise two restarts with steps // 4 steps.
|
||||
return [(steps // 4) + 1, 2, 0.1, a1111_t_max]
|
||||
|
||||
restart_info = restart_info.strip().lower()
|
||||
if restart_info == "":
|
||||
# No restarts.
|
||||
return []
|
||||
restart_arrays = None
|
||||
if restart_info == "default":
|
||||
restart_info = DEFAULT_SEGMENTS
|
||||
elif restart_info == "a1111":
|
||||
restart_arrays = [get_a1111_segment()]
|
||||
if restart_arrays == [[]]:
|
||||
return []
|
||||
if restart_arrays is None:
|
||||
try:
|
||||
restart_arrays = ast.literal_eval(f"[{restart_info}]")
|
||||
except SyntaxError:
|
||||
print("Ill-formed restart segments")
|
||||
raise
|
||||
temp = []
|
||||
default_segments = ast.literal_eval(DEFAULT_SEGMENTS)
|
||||
# This phase expands any preset strings into actual 4-item restart segments.
|
||||
for idx in range(len(restart_arrays)):
|
||||
item = restart_arrays[idx]
|
||||
if not isinstance(item, str):
|
||||
temp.append(item)
|
||||
continue
|
||||
preset = item.strip().lower()
|
||||
if preset == "default":
|
||||
temp += default_segments
|
||||
elif preset == "a1111":
|
||||
temp.append(get_a1111_segment())
|
||||
else:
|
||||
raise ValueError("Ill-formed restart segment")
|
||||
restart_arrays = temp
|
||||
restart_segments = []
|
||||
# Now we build the actual restart segments.
|
||||
for arr in restart_arrays:
|
||||
if len(arr) != 4:
|
||||
raise ValueError("Restart segment must have 4 values")
|
||||
n_restart, k, t_min, t_max = arr
|
||||
if not isinstance(arr, (list, tuple)) or len(arr) != 4:
|
||||
raise ValueError("Restart segment must be a list with 4 values")
|
||||
n_restart, k, val_min, val_max = arr
|
||||
n_restart, k = int(n_restart), int(k)
|
||||
restart_segments = add_restart_segment(restart_segments, n_restart, k, t_min, t_max)
|
||||
t_min = resolve_t_value(val_min, ms)
|
||||
t_max = resolve_t_value(val_max, ms)
|
||||
restart_segments = add_restart_segment(
|
||||
restart_segments,
|
||||
n_restart,
|
||||
k,
|
||||
t_min,
|
||||
t_max,
|
||||
)
|
||||
return restart_segments
|
||||
|
||||
|
||||
@@ -39,108 +119,629 @@ def round_restart_segments(ts, restart_segments):
|
||||
:return: dict of the form {nearest_t_min: {'n': n, 'k': k, 't_max': t_max}}
|
||||
"""
|
||||
t_min_mapping = {}
|
||||
for segment in reversed(restart_segments): # Reversed to prioritize segments to the front
|
||||
t_min_neighbor = min(ts, key=lambda ts: abs(ts - segment['t_min'])).item()
|
||||
for segment in reversed(
|
||||
restart_segments,
|
||||
): # Reversed to prioritize segments to the front
|
||||
t_min_neighbor = min(ts, key=lambda ts: abs(ts - segment["t_min"])).item()
|
||||
if t_min_neighbor == ts[0]:
|
||||
warnings.warn(
|
||||
f"\n[Restart Sampling] nearest neighbor of segment t_min {segment['t_min']:.4f} is equal to the first t_min in the denoise schedule {ts[0]:.4f}, ignoring segment...", stacklevel=2)
|
||||
f"\n[Restart Sampling] nearest neighbor of segment t_min {segment['t_min']:.4f} is equal to the first t_min in the denoise schedule {ts[0]:.4f}, ignoring segment...",
|
||||
stacklevel=2,
|
||||
)
|
||||
continue
|
||||
if t_min_neighbor > segment['t_max']:
|
||||
if t_min_neighbor > segment["t_max"]:
|
||||
warnings.warn(
|
||||
f"\n[Restart Sampling] t_min neighbor {t_min_neighbor:.4f} is greater than t_max {segment['t_max']:.4f}, ignoring segment...", stacklevel=2)
|
||||
f"\n[Restart Sampling] t_min neighbor {t_min_neighbor:.4f} is greater than t_max {segment['t_max']:.4f}, ignoring segment...",
|
||||
stacklevel=2,
|
||||
)
|
||||
continue
|
||||
if t_min_neighbor in t_min_mapping:
|
||||
warnings.warn(
|
||||
f"\n[Restart Sampling] Overwriting segment {t_min_mapping[t_min_neighbor]}, nearest neighbor of {segment['t_min']:.4f} is {t_min_neighbor:.4f}", stacklevel=2)
|
||||
t_min_mapping[t_min_neighbor] = {'n': segment['n'], 'k': segment['k'], 't_max': segment['t_max']}
|
||||
f"\n[Restart Sampling] Overwriting segment {t_min_mapping[t_min_neighbor]}, nearest neighbor of {segment['t_min']:.4f} is {t_min_neighbor:.4f}",
|
||||
stacklevel=2,
|
||||
)
|
||||
t_min_mapping[t_min_neighbor] = {
|
||||
"n": segment["n"],
|
||||
"k": segment["k"],
|
||||
"t_max": segment["t_max"],
|
||||
}
|
||||
return t_min_mapping
|
||||
|
||||
|
||||
def calc_sigmas(scheduler, n, sigma_min, sigma_max, model, device):
|
||||
return SCHEDULER_MAPPING[scheduler](model, n, sigma_min, sigma_max, device)
|
||||
def calc_sigmas(
|
||||
scheduler,
|
||||
n,
|
||||
sigma_min,
|
||||
sigma_max,
|
||||
model,
|
||||
device,
|
||||
restart_segment=True,
|
||||
):
|
||||
mapping = RESTART_SCHEDULER_MAPPING if restart_segment else NORMAL_SCHEDULER_MAPPING
|
||||
return mapping[scheduler](model, n, sigma_min, sigma_max, device)
|
||||
|
||||
|
||||
def calc_restart_steps(restart_segments):
|
||||
restart_steps = 0
|
||||
for segment in restart_segments.values():
|
||||
restart_steps += (segment['n'] - 1) * segment['k']
|
||||
return restart_steps
|
||||
|
||||
|
||||
_total_steps = 0
|
||||
_restart_segments = None
|
||||
_restart_scheduler = None
|
||||
|
||||
|
||||
def restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, restart_info, restart_scheduler, denoise=1.0, disable_noise=False, start_step=None, last_step=None, force_full_denoise=False):
|
||||
global _total_steps, _restart_segments, _restart_scheduler
|
||||
_restart_scheduler = restart_scheduler
|
||||
_restart_segments = prepare_restart_segments(restart_info)
|
||||
|
||||
if sampler_name == "ddim":
|
||||
# ddim is redirected to euler
|
||||
sampler_wrapper = KSamplerRestartWrapper("euler")
|
||||
def restart_sampling(
|
||||
model,
|
||||
seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler,
|
||||
scheduler,
|
||||
positive,
|
||||
negative,
|
||||
latent_image,
|
||||
restart_info,
|
||||
restart_scheduler,
|
||||
denoise=1.0,
|
||||
disable_noise=False,
|
||||
step_range=None,
|
||||
force_full_denoise=False,
|
||||
output_only=True,
|
||||
custom_noise=None,
|
||||
chunked_mode=True,
|
||||
sigmas=None,
|
||||
):
|
||||
if isinstance(sampler, str):
|
||||
# Only possible to determine this when the sampler is passed by name. When using
|
||||
# a custom sampler, the user will need to slice sigmas when desirable.
|
||||
discard_penultimate_sigma = sampler in getattr(
|
||||
KSampler,
|
||||
"DISCARD_PENULTIMATE_SIGMA_SAMPLERS",
|
||||
set(),
|
||||
)
|
||||
sampler = sampler_object(sampler)
|
||||
else:
|
||||
sampler_wrapper = KSamplerRestartWrapper(sampler_name)
|
||||
discard_penultimate_sigma = False
|
||||
|
||||
plan = RestartPlan(
|
||||
model,
|
||||
steps,
|
||||
scheduler,
|
||||
restart_info,
|
||||
restart_scheduler,
|
||||
denoise=denoise,
|
||||
step_range=step_range,
|
||||
force_full_denoise=force_full_denoise,
|
||||
sigmas=sigmas,
|
||||
discard_penultimate_sigma=discard_penultimate_sigma,
|
||||
)
|
||||
|
||||
if VERBOSE:
|
||||
plan.explain(chunked_mode)
|
||||
|
||||
total_steps = plan.total_steps
|
||||
sigmas = plan.sigmas().to(model.load_device)
|
||||
|
||||
latent = latent_image
|
||||
latent_image = latent["samples"]
|
||||
|
||||
if disable_noise:
|
||||
torch.manual_seed(
|
||||
seed,
|
||||
) # workaround for https://github.com/comfyanonymous/ComfyUI/issues/2833
|
||||
noise = torch.zeros(
|
||||
latent_image.size(),
|
||||
dtype=latent_image.dtype,
|
||||
layout=latent_image.layout,
|
||||
device="cpu",
|
||||
)
|
||||
else:
|
||||
batch_inds = latent.get("batch_index", None)
|
||||
noise = prepare_noise(latent_image, seed, batch_inds)
|
||||
|
||||
noise_mask = None
|
||||
if "noise_mask" in latent:
|
||||
noise_mask = latent["noise_mask"]
|
||||
|
||||
x0_output = {}
|
||||
callback = latent_preview.prepare_callback(model, plan.total_steps, x0_output)
|
||||
|
||||
disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED
|
||||
|
||||
restart_options = {
|
||||
"restart_chunked": chunked_mode,
|
||||
"restart_wrapped_sampler": sampler,
|
||||
"restart_custom_noise": custom_noise,
|
||||
}
|
||||
|
||||
ksampler = KSAMPLER(
|
||||
RestartSampler.sampler_function,
|
||||
extra_options=sampler.extra_options | restart_options,
|
||||
inpaint_options=sampler.inpaint_options | {},
|
||||
)
|
||||
|
||||
# Add the additional steps to the progress bar
|
||||
pbar_update_absolute = ProgressBar.update_absolute
|
||||
|
||||
def pbar_update_absolute_wrapper(self, value, total=None, preview=None):
|
||||
pbar_update_absolute(self, value, _total_steps, preview)
|
||||
def pbar_update_absolute_wrapper(self, value, total=None, preview=None): # noqa: ARG001
|
||||
pbar_update_absolute(self, value, total_steps, preview)
|
||||
|
||||
ProgressBar.update_absolute = pbar_update_absolute_wrapper
|
||||
|
||||
try:
|
||||
samples = common_ksampler(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise=denoise,
|
||||
disable_noise=disable_noise, start_step=start_step, last_step=last_step, force_full_denoise=force_full_denoise)
|
||||
samples = sample_custom(
|
||||
model,
|
||||
noise,
|
||||
cfg,
|
||||
ksampler,
|
||||
sigmas,
|
||||
positive,
|
||||
negative,
|
||||
latent_image,
|
||||
noise_mask=noise_mask,
|
||||
callback=callback,
|
||||
disable_pbar=disable_pbar,
|
||||
seed=seed,
|
||||
)
|
||||
finally:
|
||||
sampler_wrapper.cleanup()
|
||||
ProgressBar.update_absolute = pbar_update_absolute
|
||||
return samples
|
||||
|
||||
out = latent.copy()
|
||||
out["samples"] = samples
|
||||
|
||||
if output_only:
|
||||
return (out,)
|
||||
|
||||
if "x0" in x0_output:
|
||||
out_denoised = latent.copy()
|
||||
out_denoised["samples"] = model.model.process_latent_out(x0_output["x0"].cpu())
|
||||
else:
|
||||
out_denoised = out
|
||||
|
||||
return (out, out_denoised)
|
||||
|
||||
|
||||
class KSamplerRestartWrapper:
|
||||
# PlanItem:
|
||||
# sigmas: Sigmas for normal (outside of a restart segment) sampling. They start from after the previous PlanItem's steps
|
||||
# if there is one or simply the beginning of sampling.
|
||||
# k: This is the same as the restart segment definition. Set to 0 if there is no restart segment.
|
||||
# restart_sigmas: Sigmas for the restart segment if it exists, otherwise None.
|
||||
# Convenience properties:
|
||||
# total_steps
|
||||
# s_min, s_max: None if no restart sigmas.
|
||||
# n_restart: 0 if no restart sigmas.
|
||||
|
||||
ksampler = None
|
||||
|
||||
def __init__(self, sampler_name):
|
||||
self.sample_func_name = "sample_{}".format(sampler_name)
|
||||
KSamplerRestartWrapper.ksampler = getattr(k_diffusion_sampling, self.sample_func_name)
|
||||
setattr(k_diffusion_sampling, self.sample_func_name, self.ksampler_restart_wrapper)
|
||||
class PlanItem(
|
||||
namedtuple(
|
||||
"PlanItem",
|
||||
["sigmas", "k", "restart_sigmas"],
|
||||
defaults=[None, 0, None],
|
||||
),
|
||||
):
|
||||
__slots__ = ()
|
||||
|
||||
def cleanup(self):
|
||||
setattr(k_diffusion_sampling, self.sample_func_name, KSamplerRestartWrapper.ksampler)
|
||||
def __new__(cls, *args: list, **kwargs: dict):
|
||||
obj = super().__new__(cls, *args, **kwargs)
|
||||
obj.validate()
|
||||
return obj
|
||||
|
||||
def validate(self, threshold=1e-06):
|
||||
if len(self.sigmas) < 2:
|
||||
raise ValueError("PlanItem: invalid normal sigmas: too short")
|
||||
t = self.sigmas.sort(descending=True, stable=True)[0].unique_consecutive()
|
||||
if not torch.equal(self.sigmas, t):
|
||||
errstr = (
|
||||
f"PlanItem: invalid normal sigmas: out of order or contains duplicates: {self}",
|
||||
)
|
||||
raise ValueError(errstr)
|
||||
if self.k == 0:
|
||||
return
|
||||
if self.k < 0:
|
||||
raise ValueError("PlanItem: invalid negative k value")
|
||||
if len(self.restart_sigmas) < 2:
|
||||
raise ValueError("PlanItem: invalid restart sigmas: too short")
|
||||
if self.s_min >= self.s_max:
|
||||
raise ValueError("PlanItem: invalid min/max: min >= max")
|
||||
if self.sigmas[-1] - self.restart_sigmas[0] > threshold:
|
||||
raise ValueError(
|
||||
"PlanItem: invalid sigmas: last normal sigma >= first restart sigma",
|
||||
)
|
||||
if self.sigmas[-1] - self.restart_sigmas[-1] > threshold:
|
||||
errstr = (
|
||||
f"PlanItem: invalid sigmas: last restart sigma {self.restart_sigmas[-1]} < last normal sigma {self.sigmas[-1]}",
|
||||
)
|
||||
raise ValueError(errstr)
|
||||
t = self.restart_sigmas.sort(descending=True, stable=True)[
|
||||
0
|
||||
].unique_consecutive()
|
||||
if not torch.equal(self.restart_sigmas, t):
|
||||
errstr = (
|
||||
f"PlanItem: invalid restart sigmas: out of order or contains duplicates: {self}",
|
||||
)
|
||||
raise ValueError(errstr)
|
||||
|
||||
@property
|
||||
def total_steps(self):
|
||||
if self.k < 1:
|
||||
return len(self.sigmas) - 1
|
||||
return (len(self.sigmas) - 1) + (len(self.restart_sigmas) - 1) * self.k
|
||||
|
||||
@property
|
||||
def s_min(self):
|
||||
return None if self.k < 1 else self.restart_sigmas[-1].item()
|
||||
|
||||
@property
|
||||
def s_max(self):
|
||||
return None if self.k < 1 else self.restart_sigmas[0].item()
|
||||
|
||||
@property
|
||||
def n_restart(self):
|
||||
return 0 if self.k < 1 else len(self.restart_sigmas) - 1
|
||||
|
||||
|
||||
class RestartPlan:
|
||||
def __init__(
|
||||
self,
|
||||
model,
|
||||
steps,
|
||||
scheduler,
|
||||
restart_info,
|
||||
restart_scheduler,
|
||||
denoise=1.0,
|
||||
step_range=None,
|
||||
force_full_denoise=False,
|
||||
sigmas=None,
|
||||
discard_penultimate_sigma=False,
|
||||
):
|
||||
if (
|
||||
denoise <= 0
|
||||
or (sigmas is None and steps < 1)
|
||||
or (sigmas is not None and len(sigmas) < 2)
|
||||
):
|
||||
self.plan = []
|
||||
self.total_steps = 0
|
||||
return
|
||||
ms = model.get_model_object("model_sampling")
|
||||
|
||||
if sigmas is None:
|
||||
effective_steps = steps if denoise > 0.9999 else int(steps / denoise)
|
||||
sigmas = calc_sigmas(
|
||||
scheduler,
|
||||
effective_steps + int(discard_penultimate_sigma), # True evaluates to 1
|
||||
float(ms.sigma_min),
|
||||
float(ms.sigma_max),
|
||||
model.model,
|
||||
"cpu",
|
||||
restart_segment=False,
|
||||
)
|
||||
if discard_penultimate_sigma:
|
||||
sigmas = torch.cat((sigmas[:-2], sigmas[-1:]))
|
||||
else:
|
||||
steps = effective_steps = len(sigmas) - 1
|
||||
steps = steps if denoise > 0.9999 else int(effective_steps * denoise)
|
||||
sigmas = sigmas.clone().detach().cpu()
|
||||
if effective_steps != steps:
|
||||
sigmas = sigmas[-(steps + 1) :]
|
||||
if step_range is not None:
|
||||
start_step, last_step = step_range
|
||||
|
||||
if last_step < len(sigmas) - 1:
|
||||
sigmas = sigmas[: last_step + 1]
|
||||
if force_full_denoise:
|
||||
sigmas[-1] = 0
|
||||
|
||||
if start_step < len(sigmas) - 1:
|
||||
sigmas = sigmas[start_step:]
|
||||
|
||||
restart_segments = prepare_restart_segments(restart_info, ms, sigmas)
|
||||
self.plan, self.total_steps = self.build_plan_items(
|
||||
model.model,
|
||||
restart_segments,
|
||||
restart_scheduler,
|
||||
sigmas,
|
||||
"cpu",
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<RestartPlan: steps={self.total_steps}, plan={self.plan}>"
|
||||
|
||||
# Builds a list of PlanItems and calculates the total number of steps. See the comments for PlanItem
|
||||
# for more information about plans.
|
||||
# Returns two values: the plan and the total steps.
|
||||
@staticmethod
|
||||
@torch.no_grad()
|
||||
def ksampler_restart_wrapper(model, x, sigmas, extra_args=None, callback=None, disable=None):
|
||||
global _total_steps, _restart_segments, _restart_scheduler
|
||||
ksampler = KSamplerRestartWrapper.ksampler
|
||||
segments = round_restart_segments(sigmas, _restart_segments)
|
||||
_total_steps = len(sigmas) - 1 + calc_restart_steps(segments)
|
||||
step = 0
|
||||
def build_plan_items(
|
||||
model,
|
||||
restart_segments,
|
||||
restart_scheduler,
|
||||
sigmas,
|
||||
device,
|
||||
) -> tuple[list, int]:
|
||||
model_sigma_min = float(model.model_sampling.sigma_min)
|
||||
segments = round_restart_segments(sigmas, restart_segments)
|
||||
plan = []
|
||||
range_start = -1
|
||||
for i in range(len(sigmas) - 1):
|
||||
if range_start == -1:
|
||||
# Starting a new plan item - main sigmas start at the current index of i.
|
||||
range_start = i
|
||||
s_min = sigmas[i + 1].item()
|
||||
seg = segments.get(s_min)
|
||||
if seg is None:
|
||||
continue
|
||||
s_max, k, n_restart = seg["t_max"], seg["k"], seg["n"]
|
||||
if k < 1 or n_restart < 2:
|
||||
continue
|
||||
if s_max <= model_sigma_min:
|
||||
errstr = f"Restart: Invalid restart segment t_max {s_max:.05} <= model minimum sigma {model_sigma_min:.05}"
|
||||
raise ValueError(errstr)
|
||||
normal_sigmas = sigmas[range_start : i + 2]
|
||||
restart_sigmas = calc_sigmas(
|
||||
restart_scheduler,
|
||||
n_restart,
|
||||
max(model_sigma_min, s_min),
|
||||
s_max,
|
||||
model,
|
||||
device=device,
|
||||
)
|
||||
if normal_sigmas[-1] != 0:
|
||||
restart_sigmas = restart_sigmas[:-1]
|
||||
restart_sigmas[-1] = s_min # Force the restart segment to end at s_min.
|
||||
plan.append(PlanItem(normal_sigmas, k, restart_sigmas))
|
||||
range_start = -1
|
||||
if range_start != -1:
|
||||
# Include sigmas after the last restart segments in the plan.
|
||||
plan.append(PlanItem(sigmas[range_start:]))
|
||||
return plan, sum(pi.total_steps for pi in plan)
|
||||
|
||||
def callback_wrapper(x):
|
||||
x["i"] = step
|
||||
if callback is not None:
|
||||
callback(x)
|
||||
with trange(_total_steps, disable=disable) as pbar:
|
||||
for i in range(len(sigmas) - 1):
|
||||
x = ksampler(model, x, torch.tensor([sigmas[i], sigmas[i + 1]],
|
||||
device=x.device), extra_args, callback_wrapper, True)
|
||||
pbar.update(1)
|
||||
def sigmas(self) -> torch.Tensor:
|
||||
# Flattens a plan into sigmas. When the first normal sigma matches the last item's
|
||||
# final sigma, we strip the first normal sigma to avoid creating duplicates.
|
||||
if not self.plan or self.total_steps < 1:
|
||||
return torch.FloatTensor([])
|
||||
|
||||
def sigmas_generator():
|
||||
prev_last = None
|
||||
for pi in self.plan:
|
||||
yield pi.sigmas if prev_last != pi.sigmas[0] else pi.sigmas[1:]
|
||||
prev_last = pi.restart_sigmas[-1] if pi.k > 0 else pi.sigmas[-1]
|
||||
for _ in range(pi.k):
|
||||
yield pi.restart_sigmas
|
||||
|
||||
return torch.flatten(torch.cat(tuple(sigmas_generator())))
|
||||
|
||||
# Dumps information about the plan to the console. It uses the normal plan execute
|
||||
# logic.
|
||||
def explain(self, chunked=True):
|
||||
def pretty_sigmas(sigmas):
|
||||
return ", ".join(f"{sig:.4}" for sig in sigmas.tolist())
|
||||
|
||||
def dump_steps(step, sigmas, restart=0):
|
||||
rlabel = f"R{restart:>3}" if restart > 0 else " "
|
||||
if chunked:
|
||||
chunk_size = len(sigmas) - 2
|
||||
step += 1
|
||||
if sigmas[i + 1].item() in segments:
|
||||
seg = segments[sigmas[i + 1].item()]
|
||||
s_min, s_max, k, n_restart = sigmas[i + 1], seg['t_max'], seg['k'], seg['n']
|
||||
seg_sigmas = calc_sigmas(_restart_scheduler, n_restart, s_min,
|
||||
s_max, model, device=x.device)
|
||||
for _ in range(k):
|
||||
x += torch.randn_like(x) * (s_max ** 2 - s_min ** 2) ** 0.5
|
||||
for j in range(n_restart - 1):
|
||||
x = ksampler(model, x, torch.tensor(
|
||||
[seg_sigmas[j], seg_sigmas[j + 1]], device=x.device), extra_args, callback_wrapper, True)
|
||||
pbar.update(1)
|
||||
step += 1
|
||||
print(
|
||||
f"[{rlabel}] Step {step:>3}..{step+chunk_size:<3}: {pretty_sigmas(sigmas)}",
|
||||
)
|
||||
step += chunk_size
|
||||
return step
|
||||
for i in range(len(sigmas) - 1):
|
||||
step += 1
|
||||
print(f"[{rlabel}] Step {step:>3}: {pretty_sigmas(sigmas[i:i+2])}")
|
||||
return step
|
||||
|
||||
print(f"** Dumping restart sampling plan (total steps {self.total_steps}):")
|
||||
step = 0
|
||||
for pi in self.plan:
|
||||
step = dump_steps(step, pi.sigmas)
|
||||
for kidx in range(pi.k):
|
||||
step = dump_steps(step, pi.restart_sigmas, kidx + 1)
|
||||
print(
|
||||
"** Plan legend: [Rn] - steps for restart #n, normal sampling steps otherwise. Ranges are inclusive.",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def self_test(
|
||||
model,
|
||||
schedules=None,
|
||||
restart_schedules=None,
|
||||
segments=None,
|
||||
min_steps=2,
|
||||
max_steps=100,
|
||||
) -> None:
|
||||
if schedules is None:
|
||||
schedules = NORMAL_SCHEDULER_MAPPING.keys()
|
||||
if restart_schedules is None:
|
||||
restart_schedules = RESTART_SCHEDULER_MAPPING.keys()
|
||||
if segments is None:
|
||||
segments = ("default", "a1111")
|
||||
for schname in schedules:
|
||||
for rschname in restart_schedules:
|
||||
for tsegs in segments:
|
||||
print(
|
||||
f"--- Test: {min_steps}..{max_steps} steps, schedules {schname}/{rschname}, segments {tsegs}",
|
||||
)
|
||||
for tsteps in range(min_steps, max_steps + 1):
|
||||
label = f"** {tsteps:03}: {schname}, {rschname}, {tsegs}:"
|
||||
try:
|
||||
_plan = RestartPlan(
|
||||
model,
|
||||
tsteps,
|
||||
schname,
|
||||
tsegs,
|
||||
rschname,
|
||||
1.0,
|
||||
)
|
||||
except ValueError as err:
|
||||
print(f"{label}\n\t!! FAIL: {err}")
|
||||
raise
|
||||
continue
|
||||
print("\n|| Done test")
|
||||
|
||||
|
||||
class RestartScaleFactors(NamedTuple):
|
||||
latent_scale: float = 1.0
|
||||
noise_scale: float = 0.0
|
||||
|
||||
@classmethod
|
||||
def build(
|
||||
cls,
|
||||
sigma_from: torch.Tensor,
|
||||
sigma_to: torch.Tensor,
|
||||
*,
|
||||
is_flow: bool,
|
||||
):
|
||||
sigma_from = sigma_from.detach().cpu().item()
|
||||
sigma_to = sigma_to.detach().cpu().item()
|
||||
if sigma_from == sigma_to:
|
||||
return cls(latent_scale=1.0, noise_scale=0.0)
|
||||
if sigma_from > sigma_to:
|
||||
raise ValueError("Can't do a restart step to a lower sigma!")
|
||||
|
||||
if not is_flow:
|
||||
return cls(
|
||||
latent_scale=1.0,
|
||||
noise_scale=max(0.0, (sigma_to**2 - sigma_from**2)) ** 0.5,
|
||||
)
|
||||
|
||||
snr_to = 1.0 - sigma_to
|
||||
if snr_to <= 0:
|
||||
# It may make more sense to clamp the noise scale to 1.0.
|
||||
return cls(latent_scale=0.0, noise_scale=max(1.0, sigma_to))
|
||||
snr_from = 1.0 - sigma_from
|
||||
latent_scale = snr_to / snr_from
|
||||
noise_scale = max(0.0, (sigma_to**2) - (latent_scale * sigma_from) ** 2) ** 0.5
|
||||
return cls(latent_scale, noise_scale)
|
||||
|
||||
def add_noise(self, latent: torch.Tensor, noise: torch.Tensor) -> torch.Tensor:
|
||||
if self.latent_scale != 1.0:
|
||||
latent *= self.latent_scale
|
||||
if self.noise_scale != 0:
|
||||
noise *= self.noise_scale
|
||||
latent += noise
|
||||
return latent
|
||||
|
||||
|
||||
class RestartSampler:
|
||||
@staticmethod
|
||||
def get_segment(sigmas: torch.Tensor) -> torch.Tensor:
|
||||
# A normal segment ends when we either reach the end of the list or
|
||||
# encounter a sigma higher than the previous.
|
||||
last_sigma = sigmas[0]
|
||||
for idx in range(1, len(sigmas)):
|
||||
sigma = sigmas[idx]
|
||||
if sigma > last_sigma:
|
||||
return sigmas[:idx]
|
||||
last_sigma = sigma
|
||||
return sigmas
|
||||
|
||||
@classmethod
|
||||
def split_sigmas(cls, sigmas, *, is_flow: bool=False):
|
||||
# This function just splits the sigmas into chunks that are sorted descending.
|
||||
# If the first sigma of a chunk is > the last sigma of the previous chunk then this
|
||||
# is a restart segment: noising the restart uses s_min=prev_chunk[-1], s_max=chunk[0].
|
||||
# It's a generator that yields tuples of (RestartScaleFactors, chunk_sigmas).
|
||||
prev_seg = None
|
||||
while len(sigmas) > 1:
|
||||
seg = cls.get_segment(sigmas)
|
||||
sigmas = sigmas[len(seg) :]
|
||||
if prev_seg is not None and seg[0] > prev_seg[-1]:
|
||||
s_min, s_max = prev_seg[-1], seg[0]
|
||||
scale_factors = RestartScaleFactors.build(s_min, s_max, is_flow=is_flow)
|
||||
else:
|
||||
scale_factors = RestartScaleFactors()
|
||||
prev_seg = seg
|
||||
yield (scale_factors, seg)
|
||||
|
||||
# Some extra explanation for a couple of these arguments:
|
||||
#
|
||||
# restart_chunked:
|
||||
# When False, the sampling function is called step-by-step with only two sigmas at a time.
|
||||
# When True, the sampling function will be called with sigmas for multiple steps at a time.
|
||||
# this means either the steps up to the next restart segment (or the end of sampling) or the steps within
|
||||
# a restart segment.
|
||||
#
|
||||
# restart_custom_noise:
|
||||
# If set to None, restart noise will just use torch.randn_like (gaussian) for noise generation. Otherwise
|
||||
# this should contain a function that takes x, sigma_min, sigma_max, seed and returns a noise sampler
|
||||
# function (which takes sigma, sigma_next) and returns a noisy tensor.
|
||||
@classmethod
|
||||
@torch.no_grad()
|
||||
def sampler_function(
|
||||
cls,
|
||||
model,
|
||||
x,
|
||||
sigmas,
|
||||
*args: list,
|
||||
restart_wrapped_sampler=None,
|
||||
restart_chunked=True,
|
||||
restart_custom_noise=None,
|
||||
callback=None,
|
||||
disable=None,
|
||||
**kwargs: dict,
|
||||
) -> torch.Tensor:
|
||||
if not restart_wrapped_sampler:
|
||||
raise ValueError("RestartSampler: missing restart_sampler option!")
|
||||
|
||||
def restart_noise(x, _s_min, _s_max, _seed):
|
||||
return lambda _s, _sn: torch.randn_like(x)
|
||||
|
||||
seed = (kwargs.get("extra_args", {}) or {}).get("seed", 42)
|
||||
if restart_custom_noise is not None:
|
||||
restart_noise = restart_custom_noise
|
||||
|
||||
is_flow = isinstance(model.inner_model.inner_model.model_sampling, model_sampling.CONST)
|
||||
|
||||
sampler = restart_wrapped_sampler.sampler_function
|
||||
|
||||
chunks = tuple(cls.split_sigmas(sigmas, is_flow=is_flow))
|
||||
total_steps = sum(len(chunk) - 1 for _noise, chunk in chunks)
|
||||
step = 0
|
||||
noise_count = 0
|
||||
with trange(total_steps, disable=disable) as pbar:
|
||||
last_cb_sigma = None
|
||||
|
||||
def cb_wrapper(cb_state):
|
||||
nonlocal step, last_cb_sigma
|
||||
curr_sigma = cb_state.get("sigma")
|
||||
curr_sigma = (
|
||||
curr_sigma.item()
|
||||
if isinstance(curr_sigma, torch.Tensor)
|
||||
else curr_sigma
|
||||
)
|
||||
if last_cb_sigma is not None and curr_sigma == last_cb_sigma:
|
||||
# No change since last time we were called, so we won't track it as a step.
|
||||
return
|
||||
step += 1
|
||||
pbar.update(1)
|
||||
cb_state["i"] = step
|
||||
last_cb_sigma = curr_sigma
|
||||
if callback is not None:
|
||||
callback(cb_state)
|
||||
|
||||
def do_sample(x, sigmas):
|
||||
return sampler(
|
||||
model,
|
||||
x,
|
||||
sigmas,
|
||||
*args,
|
||||
callback=cb_wrapper,
|
||||
disable=True,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
for scale_factors, chunk_sigmas in chunks:
|
||||
if scale_factors.noise_scale != 0:
|
||||
s_min, s_max = chunk_sigmas[-1], chunk_sigmas[0]
|
||||
x = scale_factors.add_noise(
|
||||
x,
|
||||
restart_noise(
|
||||
x,
|
||||
s_min,
|
||||
s_max,
|
||||
seed + noise_count,
|
||||
)(
|
||||
s_max,
|
||||
s_min,
|
||||
),
|
||||
)
|
||||
noise_count += 1
|
||||
if restart_chunked:
|
||||
x = do_sample(x, chunk_sigmas)
|
||||
continue
|
||||
for i in range(len(chunk_sigmas) - 1):
|
||||
x = do_sample(x, chunk_sigmas[i : i + 2])
|
||||
return x
|
||||
|
||||
+79
-34
@@ -1,82 +1,127 @@
|
||||
import comfy
|
||||
import torch
|
||||
from comfy.k_diffusion import sampling as k_diffusion_sampling
|
||||
|
||||
# from comfy.samplers import normal_scheduler
|
||||
|
||||
|
||||
def get_sigmas_karras(model, n, s_min, s_max, device):
|
||||
# These two may be wrong for v-pred... but it seems to work?
|
||||
# Copied from k_diffusion
|
||||
def sigma_to_t(ms, sigma, quantize=True):
|
||||
log_sigmas = ms.log_sigmas.cpu()
|
||||
log_sigma = sigma.log()
|
||||
dists = log_sigma - log_sigmas[:, None]
|
||||
if quantize:
|
||||
return dists.abs().argmin(dim=0).view(sigma.shape)
|
||||
low_idx = dists.ge(0).cumsum(dim=0).argmax(dim=0).clamp(max=log_sigmas.shape[0] - 2)
|
||||
high_idx = low_idx + 1
|
||||
low, high = log_sigmas[low_idx], log_sigmas[high_idx]
|
||||
w = (low - log_sigma) / (low - high)
|
||||
w = w.clamp(0, 1)
|
||||
t = (1 - w) * low_idx + w * high_idx
|
||||
return t.view(sigma.shape)
|
||||
|
||||
|
||||
# Copied from k_diffusion
|
||||
def t_to_sigma(ms, t):
|
||||
t = t.float()
|
||||
low_idx, high_idx, w = t.floor().long(), t.ceil().long(), t.frac()
|
||||
log_sigma = (1 - w) * ms.log_sigmas[low_idx] + w * ms.log_sigmas[high_idx]
|
||||
return log_sigma.exp()
|
||||
|
||||
|
||||
def get_sigmas_karras(_model, n, s_min, s_max, device):
|
||||
return k_diffusion_sampling.get_sigmas_karras(n, s_min, s_max, device=device)
|
||||
|
||||
|
||||
def get_sigmas_exponential(model, n, s_min, s_max, device):
|
||||
def get_sigmas_exponential(_model, n, s_min, s_max, device):
|
||||
return k_diffusion_sampling.get_sigmas_exponential(n, s_min, s_max, device=device)
|
||||
|
||||
|
||||
def normal_scheduler(model, steps, s_min, s_max, sgm=False, floor=False):
|
||||
def normal_scheduler(model, steps, s_min, s_max, sgm=False):
|
||||
"""
|
||||
Pulled from comfy.samplers.normal_scheduler
|
||||
"""
|
||||
s = model.model_sampling
|
||||
start = s.timestep(torch.tensor(s_max))
|
||||
end = s.timestep(torch.tensor(s_min))
|
||||
ms = model.model_sampling
|
||||
start = ms.timestep(torch.tensor(s_max))
|
||||
end = ms.timestep(torch.tensor(s_min))
|
||||
|
||||
if sgm:
|
||||
timesteps = torch.linspace(start, end, steps + 1)[:-1]
|
||||
else:
|
||||
timesteps = torch.linspace(start, end, steps)
|
||||
|
||||
sigs = []
|
||||
for x in range(len(timesteps)):
|
||||
ts = timesteps[x]
|
||||
sigs.append(s.sigma(ts))
|
||||
sigs += [0.0]
|
||||
sigs = (*(ms.sigma(timesteps[x]) for x in range(len(timesteps))), 0.0)
|
||||
return torch.FloatTensor(sigs)
|
||||
|
||||
|
||||
def get_sigmas_normal(model, n, s_min, s_max, device):
|
||||
return normal_scheduler(model.inner_model.inner_model, n, s_min, s_max).to(device)
|
||||
return normal_scheduler(model, n, s_min, s_max).to(device)
|
||||
|
||||
|
||||
def get_sigmas_simple(model, n, s_min, s_max, device):
|
||||
min_idx = torch.argmin(torch.abs(model.sigmas - s_min))
|
||||
max_idx = torch.argmin(torch.abs(model.sigmas - s_max))
|
||||
sigmas_slice = model.sigmas[min_idx:max_idx]
|
||||
ms = model.model_sampling
|
||||
min_idx = torch.argmin(torch.abs(ms.sigmas - s_min))
|
||||
max_idx = torch.argmin(torch.abs(ms.sigmas - s_max))
|
||||
sigmas_slice = ms.sigmas[min_idx:max_idx]
|
||||
ss = len(sigmas_slice) / n
|
||||
sigs = [float(s_max)]
|
||||
for x in range(1, n - 1):
|
||||
sigs += [float(sigmas_slice[-(1 + int(x * ss))])]
|
||||
sigs += [float(s_min), 0.0]
|
||||
sigs = (
|
||||
float(s_max),
|
||||
*(float(sigmas_slice[-(1 + int(x * ss))]) for x in range(1, n - 1)),
|
||||
float(s_min),
|
||||
0.0,
|
||||
)
|
||||
return torch.tensor(sigs, device=device)
|
||||
|
||||
|
||||
def get_sigmas_ddim_uniform(model, n, s_min, s_max, device):
|
||||
t_min, t_max = model.sigma_to_t(torch.tensor([s_min, s_max], device=device))
|
||||
ms = model.model_sampling
|
||||
t_min, t_max = sigma_to_t(ms, torch.tensor([s_min, s_max], device=device))
|
||||
ddim_timesteps = torch.linspace(t_max, t_min, n, dtype=torch.int16, device=device)
|
||||
sigs = []
|
||||
for ts in ddim_timesteps:
|
||||
if ts > 999:
|
||||
ts = 999
|
||||
sigs.append(model.t_to_sigma(ts))
|
||||
sigs += [0.0]
|
||||
sigs = (*(t_to_sigma(ms, min(ts, 999)) for ts in ddim_timesteps), 0.0)
|
||||
return torch.tensor(sigs, device=device)
|
||||
|
||||
|
||||
def get_sigmas_sgm_uniform(model, n, s_min, s_max, device):
|
||||
return normal_scheduler(model, n, s_min, s_max, sgm=True).to(device)
|
||||
|
||||
|
||||
def get_sigmas_simple_test(model, n, s_min, s_max, device):
|
||||
min_idx = torch.argmin(torch.abs(model.sigmas - s_min))
|
||||
max_idx = torch.argmin(torch.abs(model.sigmas - s_max))
|
||||
sigmas_slice = model.sigmas[min_idx:max_idx]
|
||||
ms = model.model_sampling
|
||||
min_idx = torch.argmin(torch.abs(ms.sigmas - s_min))
|
||||
max_idx = torch.argmin(torch.abs(ms.sigmas - s_max))
|
||||
sigmas_slice = ms.sigmas[min_idx:max_idx]
|
||||
ss = len(sigmas_slice) / n
|
||||
sigs = []
|
||||
for x in range(n):
|
||||
sigs += [float(sigmas_slice[-(1 + int(x * ss))])]
|
||||
sigs += [0.0]
|
||||
sigs = (*(float(sigmas_slice[-(1 + int(x * ss))]) for x in range(n)), 0.0)
|
||||
return torch.tensor(sigs, device=device)
|
||||
|
||||
|
||||
SCHEDULER_MAPPING = {
|
||||
def get_comfy_scheduler_fn(name):
|
||||
return (
|
||||
lambda model,
|
||||
steps,
|
||||
_smin,
|
||||
_smax,
|
||||
device="cpu": comfy.samplers.calculate_sigmas(
|
||||
model.model_sampling,
|
||||
name,
|
||||
steps,
|
||||
).to(device)
|
||||
)
|
||||
|
||||
|
||||
RESTART_SCHEDULER_MAPPING = {
|
||||
"normal": get_sigmas_normal,
|
||||
"karras": get_sigmas_karras,
|
||||
"exponential": get_sigmas_exponential,
|
||||
"simple": get_sigmas_simple,
|
||||
"ddim_uniform": get_sigmas_ddim_uniform,
|
||||
"sgm_uniform": get_sigmas_sgm_uniform,
|
||||
"simple_test": get_sigmas_simple_test,
|
||||
}
|
||||
|
||||
NORMAL_SCHEDULER_MAPPING = {
|
||||
k: get_comfy_scheduler_fn(k) for k in comfy.samplers.SCHEDULER_NAMES
|
||||
} | {
|
||||
"simple_test": get_sigmas_simple_test,
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user