Rebuild nodes.py from scratch; fix Flux.2's actual 128ch/16-stride latent spec

Ground-up rewrite of nodes.py/__init__.py/web extension, keeping node IDs,
model order, and presets.py unchanged for backward compatibility.

Found the real bug while rebuilding: Flux.2 uses 128 channels at 16x spatial
downsample (comfy_extras/nodes_flux.py EmptyFlux2LatentImage), not 16ch/8x
like Flux.1/Krea/Qwen-Image. We hardcoded //8 everywhere, producing a latent
2x too large per axis for Flux.2 (4x the pixels) — explains both the wrong
output resolution and the outsized s/it gap vs the native node.

By Axis node previously had no model awareness at all, always assuming
4ch/8-stride (SD1.5-style) regardless of target model — silently wrong for
every 16/128-channel model. Added a model input so it uses the correct
per-model latent spec too. Also fixed two Flux.2 presets (1080 -> 1088) that
weren't multiples of its 16-stride requirement.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
Budi Hartono
2026-08-14 19:26:53 +07:00
co-authored by Claude Sonnet 5
parent c434c7924f
commit 8dad556fff
3 changed files with 79 additions and 64 deletions
+2 -3
View File
@@ -1,12 +1,11 @@
from .nodes import EmptyLatentAspectPreset, EmptyLatentAspectByAxis
from .nodes import EmptyLatentAspectByAxis, EmptyLatentAspectPreset
# Node IDs are referenced by users' saved workflows — never rename.
NODE_CLASS_MAPPINGS = {
"CAS Empty Latent Aspect Ratio Preset": EmptyLatentAspectPreset,
"CAS Empty Latent Aspect Ratio Axis": EmptyLatentAspectByAxis,
}
NODE_DISPLAY_NAME = "latent"
WEB_DIRECTORY = "web"
__all__ = ["NODE_CLASS_MAPPINGS", "WEB_DIRECTORY"]
+90 -74
View File
@@ -1,76 +1,31 @@
import torch
from .presets import PRESETS
try:
# Same device ComfyUI's own EmptyLatentImage/EmptySD3LatentImage create latents on,
# instead of defaulting to CPU and paying a CPU->device copy at the start of sampling.
# Same device ComfyUI's own EmptyLatentImage/EmptySD3LatentImage create latents on.
# A CPU-resident latent forces per-step CPU<->GPU transfers during sampling.
import comfy.model_management as _model_management
_LATENT_DEVICE = _model_management.intermediate_device()
except ImportError: # allows importing this module outside a running ComfyUI (e.g. sanity checks)
_LATENT_DEVICE = "cpu"
def _validate_dim(v: int):
if v <= 0 or v % 8 != 0:
raise ValueError("Dimension must be >0 and divisible by 8")
class EmptyLatentAspectPreset:
"""Creates a blank latent using one of the predefined presets."""
# Built once at import time and shared by all instances/calls.
PRESET_MAP = {
f"{w}x{h} - {lbl} - {model}": (w, h)
for model, lbl, w, h in PRESETS
# Per-model (channels, spatial_stride) latent spec, matching each model's native
# ComfyUI empty-latent node. SD1.5/SDXL/Flux.1/Krea/Qwen-Image all use an 8x-downsampled
# VAE (4ch legacy or 16ch SD3-family); Flux.2 is the outlier — comfy_extras/nodes_flux.py's
# EmptyFlux2LatentImage uses 128 channels at 16x downsample. Getting either value wrong
# silently produces a latent of the wrong shape (e.g. 2x the intended resolution per axis).
MODEL_LATENT_SPEC = {
"SD15": (4, 8),
"SDXL": (4, 8),
"Flux.1": (16, 8),
"Flux.2": (128, 16),
"Krea": (16, 8),
"Qwen-Image": (16, 8),
}
# Unique models in first-appearance order; the "model" widget below is purely a
# client-side filter (see web/aspect_ratio_filter.js) for the "preset" dropdown,
# which already encodes the model in its label — parsed there, not duplicated.
MODELS = list(dict.fromkeys(model for model, _, _, _ in PRESETS))
DEFAULT_LATENT_SPEC = (4, 8)
# Latent channel count per model family. SD1.5/SDXL use the legacy 4-channel VAE;
# Flux/Qwen-Image/Krea use a 16-channel VAE (same family as ComfyUI's own
# EmptySD3LatentImage) — feeding a 4-channel latent into those samplers fails or
# silently produces garbage.
LATENT_CHANNELS = {
"SD15": 4,
"SDXL": 4,
"Flux.1": 16,
"Flux.2": 16,
"Krea": 16,
"Qwen-Image": 16,
}
DEFAULT_LATENT_CHANNELS = 4
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": (cls.MODELS,),
"preset": (list(cls.PRESET_MAP.keys()),),
"batch_size": ("INT", {"default": 1, "min": 1})
}
}
RETURN_TYPES = ("LATENT", "INT", "INT")
RETURN_NAMES = ("LATENT", "width", "height")
FUNCTION = "generate"
CATEGORY = "latent" # moved into ComfyUI's built-in "latent" category
def generate(self, model: str, preset: str, batch_size: int):
if preset not in self.PRESET_MAP:
raise ValueError(f"Unknown preset: {preset}")
w, h = self.PRESET_MAP[preset]
_validate_dim(w)
_validate_dim(h)
channels = self.LATENT_CHANNELS.get(model, self.DEFAULT_LATENT_CHANNELS)
latent = torch.zeros([batch_size, channels, h // 8, w // 8], dtype=torch.float32, device=_LATENT_DEVICE)
return ({"samples": latent}, w, h)
class EmptyLatentAspectByAxis:
"""Creates a blank latent by fixing one axis and computing the other from an aspect ratio."""
ASPECT_CHOICES = [
ASPECT_RATIOS = [
("1:1 Square", (1, 1)),
("3:2 Landscape", (3, 2)),
("2:3 Portrait", (2, 3)),
@@ -85,39 +40,100 @@ class EmptyLatentAspectByAxis:
("7:5 Landscape", (7, 5)),
("5:7 Portrait", (5, 7)),
]
REFERENCE_CHOICES = ["Width", "Height"]
def _validate_dim(v: int, stride: int = 8):
if v <= 0 or v % stride != 0:
raise ValueError(f"Dimension must be >0 and divisible by {stride}")
def _empty_latent(batch_size: int, channels: int, stride: int, w: int, h: int):
return {"samples": torch.zeros(
[batch_size, channels, h // stride, w // stride],
dtype=torch.float32,
device=_LATENT_DEVICE,
)}
class EmptyLatentAspectPreset:
"""Creates a blank latent using one of the predefined presets."""
# Built once at import time and shared by all instances/calls.
PRESET_MAP = {
f"{w}x{h} - {lbl} - {model}": (w, h)
for model, lbl, w, h in PRESETS
}
# Unique models in first-appearance order; the "model" widget is purely a
# client-side filter (see web/aspect_ratio_filter.js) for the "preset" dropdown,
# which already encodes the model in its label — parsed there, not duplicated.
MODELS = list(dict.fromkeys(model for model, _, _, _ in PRESETS))
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"primary_dim": ("INT", {"default": 512, "min": 8}),
"reference": (cls.REFERENCE_CHOICES,),
"aspect_ratio": ([lbl for lbl,_ in cls.ASPECT_CHOICES],),
"batch_size": ("INT", {"default": 1, "min": 1})
"model": (cls.MODELS,),
"preset": (list(cls.PRESET_MAP.keys()),),
"batch_size": ("INT", {"default": 1, "min": 1}),
}
}
RETURN_TYPES = ("LATENT", "INT", "INT")
RETURN_NAMES = ("LATENT", "width", "height")
FUNCTION = "generate"
CATEGORY = "latent" # now appears under the built-in latent category
CATEGORY = "latent"
def generate(self, model: str, preset: str, batch_size: int):
if preset not in self.PRESET_MAP:
raise ValueError(f"Unknown preset: {preset}")
w, h = self.PRESET_MAP[preset]
channels, stride = MODEL_LATENT_SPEC.get(model, DEFAULT_LATENT_SPEC)
_validate_dim(w, stride)
_validate_dim(h, stride)
return (_empty_latent(batch_size, channels, stride, w, h), w, h)
class EmptyLatentAspectByAxis:
"""Creates a blank latent by fixing one axis and computing the other from an aspect ratio."""
ASPECT_CHOICES = ASPECT_RATIOS
REFERENCE_CHOICES = ["Width", "Height"]
RATIO_MAP = dict(ASPECT_CHOICES)
MODELS = list(MODEL_LATENT_SPEC.keys())
def generate(self, primary_dim: int, reference: str, aspect_ratio: str, batch_size: int):
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": (cls.MODELS,),
"primary_dim": ("INT", {"default": 512, "min": 8}),
"reference": (cls.REFERENCE_CHOICES,),
"aspect_ratio": ([lbl for lbl, _ in cls.ASPECT_CHOICES],),
"batch_size": ("INT", {"default": 1, "min": 1}),
}
}
RETURN_TYPES = ("LATENT", "INT", "INT")
RETURN_NAMES = ("LATENT", "width", "height")
FUNCTION = "generate"
CATEGORY = "latent"
def generate(self, model: str, primary_dim: int, reference: str, aspect_ratio: str, batch_size: int):
if aspect_ratio not in self.RATIO_MAP:
raise ValueError(f"Unknown aspect ratio: {aspect_ratio}")
wr, hr = self.RATIO_MAP[aspect_ratio]
_validate_dim(primary_dim)
channels, stride = MODEL_LATENT_SPEC.get(model, DEFAULT_LATENT_SPEC)
_validate_dim(primary_dim, stride)
if reference == "Width":
w, h = primary_dim, round(primary_dim * hr / wr)
else:
h, w = primary_dim, round(primary_dim * wr / hr)
_validate_dim(w)
_validate_dim(h)
_validate_dim(w, stride)
_validate_dim(h, stride)
latent = torch.zeros([batch_size, 4, h // 8, w // 8], dtype=torch.float32, device=_LATENT_DEVICE)
return ({"samples": latent}, w, h)
return (_empty_latent(batch_size, channels, stride, w, h), w, h)
+2 -2
View File
@@ -20,11 +20,11 @@ PRESETS = [
("Flux.2", "1:1 Square", 2048, 2048),
("Flux.2", "3:2 Landscape", 1728, 1152),
("Flux.2", "4:3 Landscape", 1664, 1248),
("Flux.2", "16:9 Landscape", 1920, 1080),
("Flux.2", "16:9 Landscape", 1920, 1088),
("Flux.2", "21:9 Landscape", 2176, 960),
("Flux.2", "2:3 Portrait", 1152, 1728),
("Flux.2", "3:4 Portrait", 1248, 1664),
("Flux.2", "9:16 Portrait", 1080, 1920),
("Flux.2", "9:16 Portrait", 1088, 1920),
("Flux.2", "9:21 Portrait", 960, 2176),
# --- Qwen-Image (native ~1.7MP, wide native AR set) ---