Add continuous VAE decoding

This commit is contained in:
野生の男 2026-07-19 19:54:21 +09:00
parent 06c5f7d584
commit 0c4ae6c329
3 changed files with 487 additions and 1 deletions

View File

@ -1,4 +1,5 @@
import asyncio
import contextvars
import logging
import math
from dataclasses import dataclass, field
@ -14,6 +15,7 @@ import comfy.model_patcher
import comfy.patcher_extension
import comfy.sampler_helpers
import comfy.samplers
import comfy.sd
import comfy.supported_models
from comfy_execution.progress import reset_progress_registry, set_progress_registry
from comfy_execution.utils import CurrentNodeContext, reset_current_client_id, set_current_client_id
@ -29,6 +31,7 @@ CONTINUOUS_SAMPLER_NODE_FAMILIES = {
}
_COORDINATORS = {}
_VAE_COORDINATORS = {}
_CANCEL_CHECKER = None
@ -660,3 +663,195 @@ async def sample_euler_continuous(state):
coordinator = ContinuousBatchCoordinator(model_key, state)
_COORDINATORS[model_key] = coordinator
return await coordinator.submit(state)
@dataclass
class ContinuousVAEDecodeRequest:
vae: Any
samples: torch.Tensor
max_batch_size: int
admission_delay: float
prompt_id: str | None = None
def validate(self):
if type(self.vae) is not comfy.sd.VAE:
raise ValueError("Continuous VAE decoding requires a core VAE")
if not torch.is_tensor(self.samples):
raise ValueError("Continuous VAE decoding requires a tensor latent")
if self.samples.is_nested:
raise ValueError("Continuous VAE decoding does not support nested latents")
if self.samples.layout is not torch.strided:
raise ValueError("Continuous VAE decoding requires a dense strided latent")
if self.samples.ndim == 0:
raise ValueError("Continuous VAE decoding requires a batched latent")
if self.samples.shape[0] != 1:
raise ValueError("Continuous VAE decoding requires one latent per request")
if self.vae.latent_dim == 2:
if self.samples.ndim != 4:
raise ValueError("Continuous VAE decoding requires a four-dimensional image latent")
elif self.vae.latent_dim == 3:
if self.samples.ndim != 5 or self.samples.shape[2] != 1:
raise ValueError("Continuous VAE decoding supports only single-frame image latents for three-dimensional VAEs")
else:
raise ValueError("Continuous VAE decoding does not support this VAE latent dimensionality")
if self.max_batch_size < 1:
raise ValueError("Continuous VAE decoding max batch size must be positive")
if not math.isfinite(self.admission_delay) or self.admission_delay < 0:
raise ValueError("Continuous VAE decoding admission delay must be finite and non-negative")
def key(self):
self.validate()
return (
id(self.vae),
id(self.vae.patcher),
self.vae.device,
self.vae.vae_dtype,
self.vae.output_device,
tuple(self.samples.shape[1:]),
self.samples.dtype,
self.samples.device,
self.max_batch_size,
self.admission_delay,
)
def vae_key(self):
return id(self.vae)
def clear(self):
self.vae = None
self.samples = None
@dataclass
class _QueuedVAERequest:
state: ContinuousVAEDecodeRequest
future: asyncio.Future = field(init=False)
cancelled: bool = False
def __post_init__(self):
self.future = asyncio.get_running_loop().create_future()
class ContinuousVAECoordinator:
def __init__(self, vae_key, vae):
self.vae_key = vae_key
self.vae = vae
self.pending = []
self.task = None
self.last_batch_size = None
async def submit(self, state):
request = _QueuedVAERequest(state)
self.pending.append(request)
if self.task is None:
self.task = contextvars.Context().run(asyncio.create_task, self._run())
try:
return await request.future
except asyncio.CancelledError:
request.cancelled = True
raise
@staticmethod
def _cancel(request):
request.state.clear()
if not request.future.done():
request.future.set_exception(comfy.model_management.InterruptProcessingException())
@staticmethod
def _fail(request, error):
request.state.clear()
if not request.future.done():
request.future.set_exception(error)
def _drop_cancelled(self):
remaining = []
for request in self.pending:
if request.cancelled or _is_cancelled(request.state.prompt_id):
self._cancel(request)
else:
remaining.append(request)
self.pending = remaining
def _take_group(self, group_key, max_batch_size):
group = []
remaining = []
for request in self.pending:
if request.cancelled or _is_cancelled(request.state.prompt_id):
self._cancel(request)
elif request.state.key() == group_key and len(group) < max_batch_size:
group.append(request)
else:
remaining.append(request)
self.pending = remaining
return group
def _decode(self, group):
if len(group) == 1:
images = self.vae.decode(group[0].state.samples)
else:
images = self.vae.decode(torch.cat([request.state.samples for request in group]))
if not torch.is_tensor(images) or images.ndim < 1 or images.shape[0] != len(group):
raise RuntimeError("Continuous VAE decoder returned an invalid batch shape")
return images
def _finish_group(self, group, images):
for index, request in enumerate(group):
if request.cancelled or _is_cancelled(request.state.prompt_id):
self._cancel(request)
continue
output = images if len(group) == 1 else images[index:index + 1].clone()
request.state.clear()
if not request.future.done():
request.future.set_result(output)
async def _run(self):
error = None
try:
while self.pending:
self._drop_cancelled()
if not self.pending:
continue
representative = self.pending[0]
group_key = representative.state.key()
if representative.state.admission_delay > 0:
await asyncio.sleep(representative.state.admission_delay)
self._drop_cancelled()
if not any(request is representative for request in self.pending):
continue
group = self._take_group(group_key, representative.state.max_batch_size)
if not group:
continue
if len(group) != self.last_batch_size:
logging.info("Continuous VAE batch size: %d", len(group))
self.last_batch_size = len(group)
try:
images = self._decode(group)
except Exception as decode_error:
for request in group:
self._fail(request, decode_error)
continue
self._finish_group(group, images)
except Exception as run_error:
error = run_error
finally:
if error is not None:
for request in self.pending:
self._fail(request, error)
else:
for request in self.pending:
self._cancel(request)
self.pending.clear()
self.vae = None
if _VAE_COORDINATORS.get(self.vae_key) is self:
del _VAE_COORDINATORS[self.vae_key]
self.task = None
async def decode_vae_continuous(state):
state.key()
vae_key = state.vae_key()
coordinator = _VAE_COORDINATORS.get(vae_key)
if coordinator is None:
coordinator = ContinuousVAECoordinator(vae_key, state.vae)
_VAE_COORDINATORS[vae_key] = coordinator
return await coordinator.submit(state)

View File

@ -340,6 +340,40 @@ class VAEDecode:
images = images.reshape(-1, images.shape[-3], images.shape[-2], images.shape[-1])
return (images, )
class ContinuousVAEDecode:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"samples": ("LATENT", {"tooltip": "The latent to be decoded."}),
"vae": ("VAE", {"tooltip": "The VAE model used for decoding the latent."}),
"max_batch_size": ("INT", {"default": 4, "min": 1, "max": 64}),
"admission_delay_ms": ("FLOAT", {"default": 2.0, "min": 0.0, "max": 1000.0, "step": 0.1}),
}
}
RETURN_TYPES = ("IMAGE",)
OUTPUT_TOOLTIPS = ("The decoded image.",)
FUNCTION = "decode"
CATEGORY = "model/latent"
DESCRIPTION = "Continuously batches compatible core VAE decodes from independent workflows."
async def decode(self, samples, vae, max_batch_size=4, admission_delay_ms=2.0):
execution_context = get_executing_context()
state = comfy.continuous_batching.ContinuousVAEDecodeRequest(
vae=vae,
samples=samples["samples"],
max_batch_size=max_batch_size,
admission_delay=admission_delay_ms / 1000.0,
prompt_id=execution_context.prompt_id if execution_context is not None else None,
)
images = await comfy.continuous_batching.decode_vae_continuous(state)
if len(images.shape) == 5: #Combine batches
images = images.reshape(-1, images.shape[-3], images.shape[-2], images.shape[-1])
return (images, )
class VAEDecodeTiled:
@classmethod
def INPUT_TYPES(s):
@ -2139,6 +2173,7 @@ NODE_CLASS_MAPPINGS = {
"CLIPTextEncode": CLIPTextEncode,
"CLIPSetLastLayer": CLIPSetLastLayer,
"VAEDecode": VAEDecode,
"ContinuousVAEDecode": ContinuousVAEDecode,
"VAEEncode": VAEEncode,
"VAEEncodeForInpaint": VAEEncodeForInpaint,
"VAELoader": VAELoader,
@ -2251,6 +2286,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"VAEEncodeForInpaint": "VAE Encode (for Inpainting)",
"SetLatentNoiseMask": "Set Latent Noise Mask",
"VAEDecode": "VAE Decode",
"ContinuousVAEDecode": "VAE Decode (Continuous)",
"VAEEncode": "VAE Encode",
"LatentRotate": "Rotate Latent",
"LatentFlip": "Flip Latent",

View File

@ -5,6 +5,7 @@ import pytest
import torch
import comfy.conds
import comfy.continuous_batching as continuous_batching
import comfy.model_base
import comfy.patcher_extension
import comfy.supported_models
@ -14,6 +15,7 @@ from comfy.continuous_batching import (
FAMILY_SDXL,
ContinuousBatchCoordinator,
ContinuousBatchSession,
ContinuousVAEDecodeRequest,
_cfg_branches,
_conditioning_structure,
_PreparedConditioning,
@ -22,11 +24,12 @@ from comfy.continuous_batching import (
_validate_model_extensions,
_validate_model_family,
cfg_combine,
decode_vae_continuous,
euler_step,
)
from comfy_execution.graph import DynamicPrompt
from comfy_execution.progress import ProgressRegistry, get_progress_state, reset_progress_state
from comfy_execution.utils import get_current_client_id, get_executing_context, reset_current_client_id, set_current_client_id
from comfy_execution.utils import CurrentNodeContext, get_current_client_id, get_executing_context, reset_current_client_id, set_current_client_id
class FakeState:
@ -80,6 +83,38 @@ class FakeSession:
self.reload_requests += 1
class FakeVAE:
def __init__(self, latent_dim=2, not_video=False, fail_shape=None, on_decode=None):
self.patcher = object()
self.device = torch.device("cpu")
self.vae_dtype = torch.float32
self.output_device = torch.device("cpu")
self.latent_dim = latent_dim
self.not_video = not_video
self.fail_shape = fail_shape
self.on_decode = on_decode
self.calls = []
def decode(self, samples):
self.calls.append(samples)
if self.on_decode is not None:
self.on_decode(samples)
if self.fail_shape is not None and tuple(samples.shape[2:]) == self.fail_shape:
self.fail_shape = None
raise RuntimeError("VAE decode failed")
return samples
def _vae_request(vae, samples, max_batch_size=4, admission_delay=0.01, prompt_id=None):
return ContinuousVAEDecodeRequest(
vae=vae,
samples=samples,
max_batch_size=max_batch_size,
admission_delay=admission_delay,
prompt_id=prompt_id,
)
def _model_patcher(model):
return SimpleNamespace(model=model)
@ -486,3 +521,223 @@ def test_failure_clears_requests_and_finishing_request_reloads_survivor():
asyncio.run(failure())
asyncio.run(residency())
def test_continuous_vae_single_request_passes_the_original_tensor(monkeypatch):
monkeypatch.setattr(continuous_batching.comfy.sd, "VAE", FakeVAE)
async def run():
vae = FakeVAE()
latent = torch.full((1, 4, 2, 2), 3.0)
state = _vae_request(vae, latent)
result = await decode_vae_continuous(state)
assert vae.calls == [latent]
assert result is latent
assert state.vae is None
assert state.samples is None
assert not continuous_batching._VAE_COORDINATORS
asyncio.run(run())
def test_continuous_vae_coalesces_four_requests_in_order_and_clones_outputs(monkeypatch):
monkeypatch.setattr(continuous_batching.comfy.sd, "VAE", FakeVAE)
async def run():
vae = FakeVAE()
latents = [torch.full((1, 4, 2, 2), value, dtype=torch.float32) for value in range(4)]
states = [_vae_request(vae, latent) for latent in latents]
results = await asyncio.gather(*[decode_vae_continuous(state) for state in states])
assert len(vae.calls) == 1
assert vae.calls[0].shape == (4, 4, 2, 2)
assert torch.equal(vae.calls[0][:, 0, 0, 0], torch.arange(4, dtype=torch.float32))
for index, result in enumerate(results):
assert result.shape == (1, 4, 2, 2)
assert torch.equal(result, latents[index])
assert result.data_ptr() != vae.calls[0][index:index + 1].data_ptr()
assert states[index].vae is None
assert states[index].samples is None
assert not continuous_batching._VAE_COORDINATORS
asyncio.run(run())
def test_continuous_vae_separates_shapes_dtypes_and_vae_identities(monkeypatch):
monkeypatch.setattr(continuous_batching.comfy.sd, "VAE", FakeVAE)
async def run():
vae = FakeVAE()
other_vae = FakeVAE()
requests = [
decode_vae_continuous(_vae_request(vae, torch.zeros(1, 4, 2, 2))),
decode_vae_continuous(_vae_request(vae, torch.zeros(1, 4, 3, 2))),
decode_vae_continuous(_vae_request(vae, torch.zeros(1, 4, 2, 2, dtype=torch.float16))),
decode_vae_continuous(_vae_request(vae, torch.zeros(1, 4, 2, 2, device="meta"))),
decode_vae_continuous(_vae_request(other_vae, torch.zeros(1, 4, 2, 2))),
]
await asyncio.gather(*requests)
assert len(vae.calls) == 4
assert all(call.shape[0] == 1 for call in vae.calls)
assert len(other_vae.calls) == 1
assert other_vae.calls[0].shape == (1, 4, 2, 2)
assert not continuous_batching._VAE_COORDINATORS
asyncio.run(run())
def test_continuous_vae_accepts_and_coalesces_single_frame_anima_latents(monkeypatch):
monkeypatch.setattr(continuous_batching.comfy.sd, "VAE", FakeVAE)
async def run():
vae = FakeVAE(latent_dim=3, not_video=False)
latents = [torch.full((1, 16, 1, 2, 2), value) for value in (1.0, 2.0)]
results = await asyncio.gather(*[decode_vae_continuous(_vae_request(vae, latent)) for latent in latents])
assert len(vae.calls) == 1
assert vae.calls[0].shape == (2, 16, 1, 2, 2)
assert all(torch.equal(result, latent) for result, latent in zip(results, latents))
asyncio.run(run())
def test_continuous_vae_rejects_unsupported_inputs(monkeypatch):
monkeypatch.setattr(continuous_batching.comfy.sd, "VAE", FakeVAE)
vae = FakeVAE()
with pytest.raises(ValueError, match="core VAE"):
_vae_request(object(), torch.zeros(1, 4, 2, 2)).key()
with pytest.raises(ValueError, match="batched latent"):
_vae_request(vae, torch.tensor(0.0)).key()
with pytest.raises(ValueError, match="one latent"):
_vae_request(vae, torch.zeros(2, 4, 2, 2)).key()
with pytest.raises(ValueError, match="four-dimensional"):
_vae_request(vae, torch.zeros(1, 4, 1, 2, 2)).key()
with pytest.raises(ValueError, match="nested"):
_vae_request(vae, torch.nested.nested_tensor([torch.zeros(4, 2, 2)], layout=torch.jagged)).key()
with pytest.raises(ValueError, match="dense strided"):
_vae_request(vae, torch.zeros(1, 4, 2, 2).to_sparse()).key()
with pytest.raises(ValueError, match="single-frame"):
_vae_request(FakeVAE(latent_dim=3, not_video=True), torch.zeros(1, 4, 2, 2, 2)).key()
with pytest.raises(ValueError, match="max batch"):
_vae_request(vae, torch.zeros(1, 4, 2, 2), max_batch_size=0).key()
def test_continuous_vae_cancel_after_decode_does_not_cancel_peer(monkeypatch):
monkeypatch.setattr(continuous_batching.comfy.sd, "VAE", FakeVAE)
cancelled = set()
continuous_batching.set_cancel_checker(cancelled.__contains__)
async def run():
vae = FakeVAE(on_decode=lambda samples: cancelled.add("cancelled"))
cancelled_state = _vae_request(vae, torch.zeros(1, 4, 2, 2), prompt_id="cancelled")
peer_state = _vae_request(vae, torch.ones(1, 4, 2, 2), prompt_id="peer")
results = await asyncio.gather(
decode_vae_continuous(cancelled_state),
decode_vae_continuous(peer_state),
return_exceptions=True,
)
assert isinstance(results[0], comfy.model_management.InterruptProcessingException)
assert torch.equal(results[1], torch.ones(1, 4, 2, 2))
assert vae.calls[0].shape[0] == 2
assert cancelled_state.vae is None
assert cancelled_state.samples is None
assert peer_state.vae is None
assert peer_state.samples is None
assert not continuous_batching._VAE_COORDINATORS
try:
asyncio.run(run())
finally:
continuous_batching.set_cancel_checker(None)
def test_continuous_vae_failure_cleans_group_and_continues_pending_group(monkeypatch):
monkeypatch.setattr(continuous_batching.comfy.sd, "VAE", FakeVAE)
async def run():
vae = FakeVAE(fail_shape=(2, 2))
failed = torch.zeros(1, 4, 2, 2)
survivor = torch.ones(1, 4, 3, 2)
failed_state = _vae_request(vae, failed)
survivor_state = _vae_request(vae, survivor)
results = await asyncio.gather(
decode_vae_continuous(failed_state),
decode_vae_continuous(survivor_state),
return_exceptions=True,
)
assert isinstance(results[0], RuntimeError)
assert results[1] is survivor
assert [tuple(call.shape[2:]) for call in vae.calls] == [(2, 2), (3, 2)]
assert failed_state.vae is None
assert failed_state.samples is None
assert survivor_state.vae is None
assert survivor_state.samples is None
assert not continuous_batching._VAE_COORDINATORS
recovered = torch.full((1, 4, 2, 2), 2.0)
assert await decode_vae_continuous(_vae_request(vae, recovered)) is recovered
assert not continuous_batching._VAE_COORDINATORS
asyncio.run(run())
def test_continuous_vae_batched_failure_clears_every_request_and_recovers(monkeypatch):
monkeypatch.setattr(continuous_batching.comfy.sd, "VAE", FakeVAE)
async def run():
vae = FakeVAE(fail_shape=(2, 2))
states = [_vae_request(vae, torch.full((1, 4, 2, 2), value)) for value in (1.0, 2.0)]
results = await asyncio.gather(*[decode_vae_continuous(state) for state in states], return_exceptions=True)
assert all(isinstance(result, RuntimeError) for result in results)
assert vae.calls[0].shape == (2, 4, 2, 2)
assert all(state.vae is None and state.samples is None for state in states)
assert not continuous_batching._VAE_COORDINATORS
recovered = torch.zeros(1, 4, 2, 2)
assert await decode_vae_continuous(_vae_request(vae, recovered)) is recovered
assert len(vae.calls) == 2
assert not continuous_batching._VAE_COORDINATORS
asyncio.run(run())
def test_continuous_vae_node_preserves_standard_five_dimensional_image_contract(monkeypatch):
import nodes
class FakeImageVAE(FakeVAE):
def decode(self, samples):
self.calls.append(samples)
return samples.movedim(1, -1)
monkeypatch.setattr(continuous_batching.comfy.sd, "VAE", FakeImageVAE)
async def run():
vae = FakeImageVAE(latent_dim=3, not_video=False)
images, = await nodes.ContinuousVAEDecode().decode(
samples={"samples": torch.zeros(1, 4, 1, 2, 2)},
vae=vae,
max_batch_size=4,
admission_delay_ms=0.0,
)
assert images.shape == (1, 2, 2, 4)
assert nodes.NODE_CLASS_MAPPINGS["ContinuousVAEDecode"] is nodes.ContinuousVAEDecode
asyncio.run(run())
def test_continuous_vae_coordinator_runs_without_the_submitter_context(monkeypatch):
seen = []
class ContextVAE(FakeVAE):
def decode(self, samples):
seen.append(get_executing_context())
return super().decode(samples)
monkeypatch.setattr(continuous_batching.comfy.sd, "VAE", ContextVAE)
async def run():
vae = ContextVAE()
with CurrentNodeContext("request", "decode"):
await decode_vae_continuous(_vae_request(vae, torch.zeros(1, 4, 2, 2)))
asyncio.run(run())
assert seen == [None]