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:
co-authored by
Claude Sonnet 5
parent
de9ac7d681
commit
c434c7924f
@@ -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)
|
||||||
Reference in New Issue
Block a user