Milestone: place empty latents on ComfyUI's intermediate device (fixes CPU-vs-GPU slowdown)

Stable revert point — all known correctness/perf issues (latent channel
count per model, device placement) are fixed as of this commit. If a
future rebuild goes wrong, `git reset --hard 579e130` (or this amended SHA)
returns to a known-good working state.

CPU-resident latents forced per-step CPU<->GPU transfers during sampling,
causing several-x slower s/it than the native EmptyLatentImage node even
at matching resolution/channels.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
Budi Hartono
2026-08-14 18:59:53 +07:00
co-authored by Claude Sonnet 5
parent de9ac7d681
commit c434c7924f
+10 -2
View File
@@ -1,6 +1,14 @@
import torch import torch
from .presets import PRESETS 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.
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): def _validate_dim(v: int):
if v <= 0 or v % 8 != 0: if v <= 0 or v % 8 != 0:
raise ValueError("Dimension must be >0 and divisible by 8") raise ValueError("Dimension must be >0 and divisible by 8")
@@ -56,7 +64,7 @@ class EmptyLatentAspectPreset:
_validate_dim(h) _validate_dim(h)
channels = self.LATENT_CHANNELS.get(model, self.DEFAULT_LATENT_CHANNELS) channels = self.LATENT_CHANNELS.get(model, self.DEFAULT_LATENT_CHANNELS)
latent = torch.zeros([batch_size, channels, h // 8, w // 8], dtype=torch.float32) latent = torch.zeros([batch_size, channels, h // 8, w // 8], dtype=torch.float32, device=_LATENT_DEVICE)
return ({"samples": latent}, w, h) return ({"samples": latent}, w, h)
@@ -111,5 +119,5 @@ class EmptyLatentAspectByAxis:
_validate_dim(w) _validate_dim(w)
_validate_dim(h) _validate_dim(h)
latent = torch.zeros([batch_size, 4, h // 8, w // 8], dtype=torch.float32) latent = torch.zeros([batch_size, 4, h // 8, w // 8], dtype=torch.float32, device=_LATENT_DEVICE)
return ({"samples": latent}, w, h) return ({"samples": latent}, w, h)