added dynamic shift scheduler
This commit is contained in:
@@ -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]
|
||||
@@ -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]
|
||||
Reference in New Issue
Block a user