47 Commits
Author SHA1 Message Date
ssitu 41ce3ce693 Merge pull request #26 from blepping/feat_flow_restarting
Support flow models
2025-12-05 22:42:31 -05:00
blepping 25d8be7160 Support flow models 2025-12-05 13:02:27 -07:00
ssitu 5b06f9623f Merge pull request #18 from blepping/feat_restart_sampler
Make restart a sampler, add node to generate sigmas
2024-05-07 18:02:07 -04:00
blepping 985fc1f3e8 Tiny cleanup 2024-05-04 05:18:57 -06:00
blepping e7c6ffe3c0 Update tested for ComfyUI commit in README 2024-05-04 04:49:04 -06:00
blepping b00812334b Discard penultimate sigma per default behavior when sampler is passed by name 2024-05-04 04:06:58 -06:00
blepping 9d65ad1b68 Add a note about setting denoise with RestartScheduler 2024-05-04 04:05:48 -06:00
blepping d17907e09d Update PlanItem comment for s_min, s_max properties, add n_restart property for consistency 2024-05-03 14:01:12 -06:00
blepping d12cb9e96b Use INCLUDE_SELFTEST variable 2024-04-30 13:13:22 -06:00
blepping 725d3aeb58 Update documentation for restart sampler/scheduler changes
Fix denoise/step range calculation for really reals this time, I hope

Use the normal scheduler list for non-restart schedulers in nodes

Revert limiting s_max to model sigma_max

Add sgm_uniform to restart schedulers list (just normal with sgm=True)
2024-04-23 13:30:12 -06:00
blepping 83470d59e0 Fix plan denoise/start/end step calculation, hopefully 2024-04-21 16:30:32 -06:00
blepping 7ef2f56798 Fix total steps calculation in actual restart sampler 2024-04-21 05:12:30 -06:00
blepping a36442dd4f Cleanups, add some comments 2024-04-21 04:56:59 -06:00
blepping 1b65c60e0c Simplify converting plan to sigmas 2024-04-21 03:55:52 -06:00
blepping f5c7a877e5 Refactor and simplify 2024-04-19 05:58:05 -06:00
blepping 471b5972c9 only generate restart segment down to model sigma_min
merge restart segments more aggressively
2024-04-18 10:56:55 -06:00
blepping 20dc195ca9 Force last sigma in restart segment to t_min + other stuff 2024-04-17 14:05:13 -06:00
blepping bda63be15c The struggle! 2024-04-17 09:35:44 -06:00
blepping e2dcfd4091 Refactor plan recovery, make plan sampling function reusable 2024-04-16 18:59:07 -06:00
blepping 2669b01670 Correctly recover s_min, hopefully 2024-04-15 14:25:08 -06:00
blepping 3752a2daf9 Make restart a sampler, add node to generate sigmas: phase 1 2024-04-15 10:35:00 -06:00
ssitu ea79890408 Merge pull request #15 from blepping/fix_step_tracking
Only increment steps/delegate to callback when the sigma changes
2024-04-07 23:08:27 -04:00
blepping a73142188a Only increment steps/delegate to callback when the sigma changes 2024-04-03 20:03:39 -06:00
ssitandblepping bb292f8ec1 Made code more concise in simple_test
Co-Authored-By: blepping <157360029+blepping@users.noreply.github.com>
2024-04-02 00:43:33 -04:00
ssitu bc784ec621 Merge pull request #14 from blepping/code_formatting
Code formatting + clean up some lints
2024-04-01 16:56:36 -04:00
blepping ae32ce995a Code formatting + clean up some lints 2024-04-01 12:40:09 -06:00
ssitu eca26065e4 Merge pull request #13 from blepping/restart_custom_noise
Allow chunked restart sampling
2024-04-01 12:22:20 -04:00
blepping 9c65214933 Attempt to get a1111 mode t_max calculation correct
Improve sigmas output when dumping restart plan
2024-04-01 03:34:24 -06:00
blepping f890b8bf62 Use correct value for a1111 mode t_max 2024-03-31 16:04:15 -06:00
blepping 17193ecdbf Minor cleanups + add a few more comments
Allow passing sigmas to main restart_sampling function
2024-03-28 02:00:34 -06:00
blepping e2e25e7a28 Add some comments describing what's going on in plan generation and execution
Allow specifying segments "default" to use the default segments

Allow specifying segments "a1111" to calculate segments like A1111

Allow setting environment variable COMFYUI_VERBOSE_RESTART_SAMPLING=1 to get some debug info

Add chunked_mode to samplers (except for simple which will use the default of True)

Documentation updates
2024-03-27 05:45:06 -06:00
blepping bbbddbd7cb Remove unused noise_multiplier param 2024-03-25 06:00:09 -06:00
blepping 67d4b62235 Refactor plan handling 2024-03-25 05:45:48 -06:00
blepping 12c0ca1044 Simplify total steps logic in wrapper 2024-03-22 14:51:20 -06:00
blepping b172908ac7 Implement chunked restart sampling 2024-03-22 14:45:04 -06:00
blepping 33ba61bb78 Make custom restart sampler with custom noise a separate node 2024-03-22 11:47:08 -06:00
blepping 22f4ee04e5 Hack to allow setting custom noise 2024-03-22 04:20:41 -06:00
ssitu b5b8999f5d Add build of comfyui 2024-03-21 00:04:32 -04:00
ssitu 39c5f85d80 Merge pull request #11 from blepping/feat_custom_restart_sampler
Add KRestartSamplerCustom + other cleanups
2024-03-20 11:38:54 -04:00
blepping b4d2fe9661 Update README with example and t_min/max format info 2024-03-17 16:19:54 -06:00
blepping b5ef71262f Remove redundant seg assignment in restart wrapper 2024-03-17 13:17:25 -06:00
blepping eebbd9cf21 Allow specify t_min/max as timesteps and percentages in Restart segments 2024-03-17 13:11:00 -06:00
blepping 1fb1b03261 Fix setting denoise 2024-03-16 03:49:52 -06:00
blepping 20efcba391 Fix assuming first/last steps would be set for non-custom restart samplers 2024-03-15 19:20:58 -06:00
blepping 6ddbaf02f4 Add KRestartSamplerCustom + other cleanups 2024-03-14 05:41:27 -06:00
ssit 24c70dce64 Remove low quality comparisons 2023-11-30 16:49:31 -05:00
ssitu f30cb0ef9c Merge pull request #7 from ssitu/ddim_removed
Update for the recent changes to DDIM.
2023-11-02 23:37:26 -04:00
10 changed files with 1204 additions and 203 deletions
+68 -16
View File
@@ -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 | ![image](https://github.com/ssitu/ComfyUI_restart_sampling/assets/57548627/7696da21-ea8c-4263-91a9-658d0f87dc47) | 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 | ![image](https://github.com/ssitu/ComfyUI_restart_sampling/assets/57548627/7696da21-ea8c-4263-91a9-658d0f87dc47) | 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 |
| --- | --- |
| ![image](./examples/heun_edm_00002_.png) | ![image](./examples/heun_edm_restarts_00002_.png) |
| ![image](./examples/heun_edm_00003_.png) | ![image](./examples/heun_edm_restarts_00003_.png) |
| ![image](./examples/heun_edm_00001_.png) | ![image](./examples/heun_edm_restarts_00001_.png) |
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

+370 -67
View File
@@ -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
View File
@@ -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
View File
@@ -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,
}