Files
ComfyUI/tests-unit/comfy_test/vae_channel_trim_test.py
christian-byrne 3bd20204f5 Warn when VAE.encode drops input channels
VAE.encode silently trims anything past output_channels, so an RGBA image
loses its alpha with no log line. Warn once per VAE and point at
SplitImageWithAlpha. No behavior change.
2026-08-06 19:11:42 -07:00

43 lines
1.1 KiB
Python

import logging
import torch
from comfy.cli_args import args
if not torch.cuda.is_available():
args.cpu = True
from comfy.sd import VAE # noqa: E402
def make_vae():
vae = VAE(sd={})
vae.crop_input = False
return vae
def test_extra_channels_are_trimmed_and_warned_once(caplog):
vae = make_vae()
pixels = torch.zeros(1, 16, 16, 4)
with caplog.at_level(logging.WARNING):
first = vae.vae_encode_crop_pixels(pixels)
second = vae.vae_encode_crop_pixels(pixels)
assert first.shape == (1, 16, 16, 3)
assert second.shape == (1, 16, 16, 3)
trim_warnings = [r for r in caplog.records if "VAE encode" in r.getMessage()]
assert len(trim_warnings) == 1
assert "SplitImageWithAlpha" in trim_warnings[0].getMessage()
def test_matching_channels_do_not_warn(caplog):
vae = make_vae()
with caplog.at_level(logging.WARNING):
out = vae.vae_encode_crop_pixels(torch.zeros(1, 16, 16, 3))
assert out.shape == (1, 16, 16, 3)
assert not [r for r in caplog.records if "VAE encode" in r.getMessage()]