diff --git a/scheduling/dynamic_shift_scheduler.py b/scheduling/dynamic_shift_scheduler.py new file mode 100644 index 0000000..635f343 --- /dev/null +++ b/scheduling/dynamic_shift_scheduler.py @@ -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] diff --git a/scheduling/sequence_length_calculator.py b/scheduling/sequence_length_calculator.py new file mode 100644 index 0000000..3cc1a8a --- /dev/null +++ b/scheduling/sequence_length_calculator.py @@ -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]