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:
co-authored by
Claude Sonnet 5
parent
c434c7924f
commit
8dad556fff
+2
-3
@@ -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 = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
"CAS Empty Latent Aspect Ratio Preset": EmptyLatentAspectPreset,
|
"CAS Empty Latent Aspect Ratio Preset": EmptyLatentAspectPreset,
|
||||||
"CAS Empty Latent Aspect Ratio Axis": EmptyLatentAspectByAxis,
|
"CAS Empty Latent Aspect Ratio Axis": EmptyLatentAspectByAxis,
|
||||||
}
|
}
|
||||||
|
|
||||||
NODE_DISPLAY_NAME = "latent"
|
|
||||||
|
|
||||||
WEB_DIRECTORY = "web"
|
WEB_DIRECTORY = "web"
|
||||||
|
|
||||||
__all__ = ["NODE_CLASS_MAPPINGS", "WEB_DIRECTORY"]
|
__all__ = ["NODE_CLASS_MAPPINGS", "WEB_DIRECTORY"]
|
||||||
@@ -1,76 +1,31 @@
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from .presets import PRESETS
|
from .presets import PRESETS
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Same device ComfyUI's own EmptyLatentImage/EmptySD3LatentImage create latents on,
|
# 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.
|
# A CPU-resident latent forces per-step CPU<->GPU transfers during sampling.
|
||||||
import comfy.model_management as _model_management
|
import comfy.model_management as _model_management
|
||||||
_LATENT_DEVICE = _model_management.intermediate_device()
|
_LATENT_DEVICE = _model_management.intermediate_device()
|
||||||
except ImportError: # allows importing this module outside a running ComfyUI (e.g. sanity checks)
|
except ImportError: # allows importing this module outside a running ComfyUI (e.g. sanity checks)
|
||||||
_LATENT_DEVICE = "cpu"
|
_LATENT_DEVICE = "cpu"
|
||||||
|
|
||||||
def _validate_dim(v: int):
|
# Per-model (channels, spatial_stride) latent spec, matching each model's native
|
||||||
if v <= 0 or v % 8 != 0:
|
# ComfyUI empty-latent node. SD1.5/SDXL/Flux.1/Krea/Qwen-Image all use an 8x-downsampled
|
||||||
raise ValueError("Dimension must be >0 and divisible by 8")
|
# 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).
|
||||||
class EmptyLatentAspectPreset:
|
MODEL_LATENT_SPEC = {
|
||||||
"""Creates a blank latent using one of the predefined presets."""
|
"SD15": (4, 8),
|
||||||
# Built once at import time and shared by all instances/calls.
|
"SDXL": (4, 8),
|
||||||
PRESET_MAP = {
|
"Flux.1": (16, 8),
|
||||||
f"{w}x{h} - {lbl} - {model}": (w, h)
|
"Flux.2": (128, 16),
|
||||||
for model, lbl, w, h in PRESETS
|
"Krea": (16, 8),
|
||||||
|
"Qwen-Image": (16, 8),
|
||||||
}
|
}
|
||||||
# Unique models in first-appearance order; the "model" widget below is purely a
|
DEFAULT_LATENT_SPEC = (4, 8)
|
||||||
# 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))
|
|
||||||
|
|
||||||
# Latent channel count per model family. SD1.5/SDXL use the legacy 4-channel VAE;
|
ASPECT_RATIOS = [
|
||||||
# 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 = [
|
|
||||||
("1:1 Square", (1, 1)),
|
("1:1 Square", (1, 1)),
|
||||||
("3:2 Landscape", (3, 2)),
|
("3:2 Landscape", (3, 2)),
|
||||||
("2:3 Portrait", (2, 3)),
|
("2:3 Portrait", (2, 3)),
|
||||||
@@ -85,39 +40,100 @@ class EmptyLatentAspectByAxis:
|
|||||||
("7:5 Landscape", (7, 5)),
|
("7:5 Landscape", (7, 5)),
|
||||||
("5:7 Portrait", (5, 7)),
|
("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
|
@classmethod
|
||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"primary_dim": ("INT", {"default": 512, "min": 8}),
|
"model": (cls.MODELS,),
|
||||||
"reference": (cls.REFERENCE_CHOICES,),
|
"preset": (list(cls.PRESET_MAP.keys()),),
|
||||||
"aspect_ratio": ([lbl for lbl,_ in cls.ASPECT_CHOICES],),
|
"batch_size": ("INT", {"default": 1, "min": 1}),
|
||||||
"batch_size": ("INT", {"default": 1, "min": 1})
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("LATENT", "INT", "INT")
|
RETURN_TYPES = ("LATENT", "INT", "INT")
|
||||||
RETURN_NAMES = ("LATENT", "width", "height")
|
RETURN_NAMES = ("LATENT", "width", "height")
|
||||||
FUNCTION = "generate"
|
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)
|
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:
|
if aspect_ratio not in self.RATIO_MAP:
|
||||||
raise ValueError(f"Unknown aspect ratio: {aspect_ratio}")
|
raise ValueError(f"Unknown aspect ratio: {aspect_ratio}")
|
||||||
wr, hr = self.RATIO_MAP[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":
|
if reference == "Width":
|
||||||
w, h = primary_dim, round(primary_dim * hr / wr)
|
w, h = primary_dim, round(primary_dim * hr / wr)
|
||||||
else:
|
else:
|
||||||
h, w = primary_dim, round(primary_dim * wr / hr)
|
h, w = primary_dim, round(primary_dim * wr / hr)
|
||||||
|
|
||||||
_validate_dim(w)
|
_validate_dim(w, stride)
|
||||||
_validate_dim(h)
|
_validate_dim(h, stride)
|
||||||
|
|
||||||
latent = torch.zeros([batch_size, 4, h // 8, w // 8], dtype=torch.float32, device=_LATENT_DEVICE)
|
return (_empty_latent(batch_size, channels, stride, w, h), w, h)
|
||||||
return ({"samples": latent}, w, h)
|
|
||||||
|
|||||||
+2
-2
@@ -20,11 +20,11 @@ PRESETS = [
|
|||||||
("Flux.2", "1:1 Square", 2048, 2048),
|
("Flux.2", "1:1 Square", 2048, 2048),
|
||||||
("Flux.2", "3:2 Landscape", 1728, 1152),
|
("Flux.2", "3:2 Landscape", 1728, 1152),
|
||||||
("Flux.2", "4:3 Landscape", 1664, 1248),
|
("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", "21:9 Landscape", 2176, 960),
|
||||||
("Flux.2", "2:3 Portrait", 1152, 1728),
|
("Flux.2", "2:3 Portrait", 1152, 1728),
|
||||||
("Flux.2", "3:4 Portrait", 1248, 1664),
|
("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),
|
("Flux.2", "9:21 Portrait", 960, 2176),
|
||||||
|
|
||||||
# --- Qwen-Image (native ~1.7MP, wide native AR set) ---
|
# --- Qwen-Image (native ~1.7MP, wide native AR set) ---
|
||||||
|
|||||||
Reference in New Issue
Block a user