added dynamic shift scheduler

This commit is contained in:
Sonny Box
2026-09-21 21:51:07 -07:00
parent 3d28d0831a
commit 8b6cbc9e6c
2 changed files with 360 additions and 0 deletions
+176
View File
@@ -0,0 +1,176 @@
import math
from comfy_api.latest import io
import torch
TIME_SHIFT_TYPES = ["exponential", "linear"]
def _calculate_mu(
seq_length: int,
base_seq_length: int,
max_seq_length: int,
base_shift: float,
max_shift: float,
) -> float:
# diffusers calculate_shift(): the straight line through
# (base_seq_length, base_shift) and (max_seq_length, max_shift), read at
# seq_length.
m = (max_shift - base_shift) / (max_seq_length - base_seq_length)
b = base_shift - m * base_seq_length
return seq_length * m + b
def _time_shift(
mu: float, t: torch.Tensor, time_shift_type: str
) -> torch.Tensor:
# diffusers _time_shift_exponential / _time_shift_linear with sigma fixed at 1.0.
# The exponential branch is the same curve as ComfyUI's shift with shift = e**mu.
if time_shift_type == "exponential":
return math.exp(mu) / (math.exp(mu) + (1.0 / t - 1.0))
return mu / (mu + (1.0 / t - 1.0))
def _stretch_to_terminal(
sigmas: torch.Tensor, shift_terminal: float
) -> torch.Tensor:
# diffusers stretch_shift_to_terminal(): rescales the schedule in (1 - sigma) space
# so the last sigma lands on shift_terminal while the first stays pinned at 1.0.
one_minus_z = 1.0 - sigmas
scale_factor = one_minus_z[-1] / (1.0 - shift_terminal)
return 1.0 - (one_minus_z / scale_factor)
class DynamicShiftScheduler(io.ComfyNode):
@classmethod
def define_schema(cls) -> io.Schema:
return io.Schema(
node_id="DynamicShiftScheduler",
display_name="🐧 Dynamic Shift Scheduler",
category="SuperNodes/Scheduling",
description="Builds a flow match sigma schedule from a diffusers scheduler config.",
search_aliases=[
"flow match",
"FlowMatchEulerDiscreteScheduler",
"dynamic shifting",
"shift terminal",
"mu",
"diffusers",
"huggingface",
],
inputs=[
io.Int.Input(
"steps",
default=20,
min=2,
max=10_000,
step=1,
),
io.Int.Input(
"seq_length",
default=4096,
min=1,
max=1_000_000,
step=1,
tooltip="Token count of the latent being sampled.",
),
io.Int.Input(
"base_seq_length",
default=256,
min=1,
max=1_000_000,
step=1,
),
io.Int.Input(
"max_seq_length",
default=4096,
min=1,
max=1_000_000,
step=1,
),
io.Float.Input(
"base_shift",
default=0.5,
min=-100.0,
max=100.0,
step=0.01,
),
io.Float.Input(
"max_shift",
default=1.15,
min=-100.0,
max=100.0,
step=0.01,
),
io.Float.Input(
"shift_terminal",
default=0.0,
min=0.0,
max=0.999,
step=0.001,
tooltip="0.0 disables it, matching a config with no shift_terminal.",
),
io.Combo.Input(
"time_shift_type",
options=TIME_SHIFT_TYPES,
),
],
outputs=[
io.Custom("SIGMAS").Output(
tooltip="The flow match sigma schedule."
),
],
)
@classmethod
def execute(
cls,
steps,
seq_length,
base_seq_length,
max_seq_length,
base_shift,
max_shift,
shift_terminal,
time_shift_type,
) -> io.NodeOutput:
if base_seq_length == max_seq_length:
raise ValueError(
f"Invalid config: base_seq_length and max_seq_length are both {base_seq_length}. "
f"The shift line needs two distinct sequence lengths to be defined."
)
mu = _calculate_mu(
seq_length,
base_seq_length,
max_seq_length,
base_shift,
max_shift,
)
# The linear form divides by mu directly rather than by e**mu, so it has no
# useful branch at or below zero the way the exponential form does
if time_shift_type == "linear" and mu <= 0.0:
raise ValueError(
f"Invalid config: linear time_shift_type produced mu = {mu:.4f} at a sequence "
f"length of {seq_length}, but it requires a positive mu. Check base_shift "
f"and max_shift, or switch to exponential."
)
# This grid is what ComfyUI's
# simple scheduler walks, so a config with no shift_terminal comes out
# the same as ModelSamplingFlux plus BasicScheduler on simple.
sigmas = torch.linspace(1.0, 1.0 / steps, steps, dtype=torch.float64)
sigmas = _time_shift(mu, sigmas, time_shift_type)
# A config without shift_terminal leaves the schedule ending on its natural
# final sigma.
if shift_terminal > 0.0:
sigmas = _stretch_to_terminal(sigmas, shift_terminal)
sigmas = torch.cat([sigmas, torch.zeros(1, dtype=torch.float64)])
return io.NodeOutput(sigmas.to(dtype=torch.float32))
NODE = [DynamicShiftScheduler]
+184
View File
@@ -0,0 +1,184 @@
import math
from comfy_api.latest import io
# (spatial compression, patch size, temporal compression)
FAMILIES = {
"Flux": (8, 2, 1), # Flux 1/2, Qwen-Image 1 and 2.1, Krea 2
# LTX-Video and LTX2: 32x spatial, 8x temporal, no patchify. 2.5 keeps the
# same geometry but turned dynamic shifting off in the diffusers config
"LTX2": (32, 1, 8),
}
MANUAL = "Manual"
def _token_count(
width: int,
height: int,
length: int,
spatial: int,
patch: int,
temporal: int,
) -> int:
# The VAE compresses first, then the transformer patchifies what is left.
# Both round up, matching pad_to_patch_size and the (h + patch // 2) // patch
# rounding the models use, so odd sizes count the padded token rather than
# dropping it. Temporal patch size is 1 on every model that reaches here.
tokens_w = math.ceil(math.ceil(width / spatial) / patch)
tokens_h = math.ceil(math.ceil(height / spatial) / patch)
return math.ceil(length / temporal) * tokens_w * tokens_h
class SequenceLengthCalculator(io.ComfyNode):
@classmethod
def define_schema(cls) -> io.Schema:
return io.Schema(
node_id="SequenceLengthCalculator",
display_name="🐧 Sequence Length Calculator",
category="SuperNodes/Scheduling",
description="Works out the token count to feed a dynamic shift scheduler.",
search_aliases=[
"sequence length",
"seq len",
"tokens",
"image_seq_len",
"dynamic shift",
"latent tokens",
],
inputs=[
io.DynamicCombo.Input(
"model_type",
tooltip="Flux is any image model at 16 px per token. LTX2 is 32 px per token with 8 frames per latent frame.",
options=[
io.DynamicCombo.Option(
"Flux",
[
io.Int.Input(
"width",
default=1024,
min=16,
max=16_384,
step=16,
),
io.Int.Input(
"height",
default=1024,
min=16,
max=16_384,
step=16,
),
],
),
io.DynamicCombo.Option(
"LTX2",
[
io.Int.Input(
"width",
default=768,
min=32,
max=16_384,
step=32,
),
io.Int.Input(
"height",
default=512,
min=32,
max=16_384,
step=32,
),
io.Int.Input(
"length",
default=97,
min=1,
max=16_384,
step=8,
),
],
),
io.DynamicCombo.Option(
MANUAL,
[
io.Int.Input(
"width",
default=1024,
min=8,
max=16_384,
step=8,
),
io.Int.Input(
"height",
default=1024,
min=8,
max=16_384,
step=8,
),
io.Int.Input(
"length",
default=1,
min=1,
max=16_384,
step=1,
tooltip="Frame count. 1 for image models.",
),
io.Int.Input(
"spatial_compression",
default=8,
min=1,
max=256,
step=1,
tooltip="How many pixels per latent pixel the VAE folds away.",
),
io.Int.Input(
"patch_size",
default=2,
min=1,
max=16,
step=1,
tooltip="How many latent pixels per side the transformer folds into one token.",
),
io.Int.Input(
"temporal_compression",
default=1,
min=1,
max=64,
step=1,
tooltip="How many frames per latent frame the VAE folds away.",
),
],
),
],
),
],
outputs=[
io.Int.Output(
display_name="seq_length",
tooltip="Token count of the latent.",
),
],
)
@classmethod
def execute(cls, model_type) -> io.NodeOutput:
family = model_type["model_type"]
if family == MANUAL:
spatial = model_type["spatial_compression"]
patch = model_type["patch_size"]
temporal = model_type["temporal_compression"]
else:
spatial, patch, temporal = FAMILIES[family]
return io.NodeOutput(
_token_count(
model_type["width"],
model_type["height"],
# Only the video and manual options carry a frame count
model_type.get("length", 1),
spatial,
patch,
temporal,
)
)
NODE = [SequenceLengthCalculator]