Merge pull request #173 from larsupb/main

tiles batch processing
This commit is contained in:
ssitu
2026-02-12 23:55:47 -05:00
committed by GitHub
10 changed files with 619 additions and 178 deletions
+2
View File
@@ -50,6 +50,7 @@ class StableDiffusionProcessing:
seam_fix_mode,
custom_sampler=None,
custom_sigmas=None,
batch_size=1,
):
# Variables used by the USDU script
self.init_images = [init_img]
@@ -85,6 +86,7 @@ class StableDiffusionProcessing:
self.upscale_by = upscale_by
self.uniform_tile_mode = uniform_tile_mode
self.tiled_decode = tiled_decode
self.batch_size = batch_size
self.vae_decoder = VAEDecode()
self.vae_encoder = VAEEncode()
self.vae_decoder_tiled = VAEDecodeTiled()
+6
View File
@@ -18,3 +18,9 @@ def save_image(tensor, path: pathlib.Path):
def load_image(path: pathlib.Path, device=None):
"""Load an image from disk and convert it to a tensor."""
return usdu_utils.pil_to_tensor(Image.open(path)).to(device=device)
def image_name_format(name: str, extension: str, batch_size: int = 1) -> str:
"""Helper for building image names for tests."""
batch_suffix = f"_batch{batch_size}" if batch_size > 1 else ""
return f"{name}{batch_suffix}{extension}"
+2 -2
View File
@@ -27,5 +27,5 @@ def test_base_image_matches_reference(base_image, test_dirs: DirectoryConfig):
diff1 = img_tensor_mae(blur(im1), blur(test_im1))
diff2 = img_tensor_mae(blur(im2), blur(test_im2))
logger.info(f"Base Image Diff1: {diff1}, Diff2: {diff2}")
assert diff1 < 0.05, "Image 1 does not match its test image."
assert diff2 < 0.05, "Image 2 does not match its test image."
assert diff1 < 0.01, "Image 1 does not match its test image."
assert diff2 < 0.01, "Image 2 does not match its test image."
+12 -18
View File
@@ -9,27 +9,26 @@ import torch
from setup_utils import execute
from tensor_utils import img_tensor_mae, blur
from io_utils import save_image, load_image
from io_utils import save_image, load_image, image_name_format
from configs import DirectoryConfig
from fixtures_images import EXT
CATEGORY = pathlib.Path(pathlib.Path(__file__).stem.removeprefix("test_"))
CONTROLNET_TILE_OUTPUT_IMAGE = "controlnet_tile" + EXT
TEST_CONTROLNET_TILE_MODEL = "control_v11f1e_sd15_tile.pth"
@pytest.mark.parametrize("batch_size", [1, 2])
class TestControlNet:
"""Integration tests for the upscaling workflow with ControlNet."""
@pytest.fixture(scope="class")
def controlnet_upscaled_image(
def test_controlnet_tile(
self,
base_image,
loaded_checkpoint,
upscale_model,
node_classes,
seed,
test_dirs,
batch_size,
test_dirs: DirectoryConfig,
):
"""Generate upscaled images using ControlNet."""
image, positive, negative = base_image
@@ -72,25 +71,20 @@ class TestControlNet:
seam_fix_padding=16,
force_uniform_tiles=True,
tiled_decode=False,
batch_size=batch_size,
)
# Save and reload sample image
sample_dir = test_dirs.sample_images
filename = CATEGORY / CONTROLNET_TILE_OUTPUT_IMAGE
filename = CATEGORY / image_name_format("controlnet_tile", EXT, batch_size)
save_image(upscaled[0], sample_dir / filename)
upscaled = load_image(sample_dir / filename)
return upscaled
def test_controlnet_upscaled_image_matches_reference(
self, controlnet_upscaled_image, test_dirs: DirectoryConfig
):
"""
Verify ControlNet upscaled images match reference images.
"""
logger = logging.getLogger("test_controlnet_upscaled_image_matches_reference")
# Verify against reference image
logger = logging.getLogger("test_controlnet_tile")
test_img_dir = test_dirs.test_images
test_img = load_image(test_img_dir / CATEGORY / CONTROLNET_TILE_OUTPUT_IMAGE)
test_img = load_image(test_img_dir / filename)
# Reduce high-frequency noise differences with gaussian blur
diff = img_tensor_mae(blur(controlnet_upscaled_image), blur(test_img))
diff = img_tensor_mae(blur(upscaled), blur(test_img))
logger.info(f"ControlNet Upscaled Image Diff: {diff}")
assert diff < 0.05, "ControlNet upscaled image does not match its test image."
assert diff < 0.01, "ControlNet upscaled image does not match its test image."
+47 -41
View File
@@ -4,24 +4,20 @@ Tests a common workflow for UltimateSDUpscale.
import logging
import pathlib
import pytest
import torch
from setup_utils import execute
from tensor_utils import img_tensor_mae, blur
from io_utils import save_image, load_image
from io_utils import save_image, load_image, image_name_format
from configs import DirectoryConfig
from fixtures_images import EXT
# Image file names
CATEGORY = pathlib.Path(pathlib.Path(__file__).stem.removeprefix("test_"))
IMAGE_1 = CATEGORY / ("main1_sd15_upscaled" + EXT)
IMAGE_2 = CATEGORY / ("main2_sd15_upscaled" + EXT)
NO_UPSCALE_IMAGE_1 = CATEGORY / ("main1_sd15_upscaled_no_upscale" + EXT)
NO_UPSCALE_IMAGE_2 = CATEGORY / ("main2_sd15_upscaled_no_upscale" + EXT)
CUSTOM_SAMPLER_IMAGE_1 = CATEGORY / ("main1_sd15_upscaled_custom_sampler" + EXT)
CUSTOM_SAMPLER_IMAGE_2 = CATEGORY / ("main2_sd15_upscaled_custom_sampler" + EXT)
@pytest.mark.parametrize("batch_size", [1, 2])
class TestMainWorkflow:
"""Integration tests for the main upscaling workflow."""
@@ -32,6 +28,7 @@ class TestMainWorkflow:
upscale_model,
node_classes,
seed,
batch_size,
test_dirs: DirectoryConfig,
):
"""Generate upscaled images using standard workflow."""
@@ -48,11 +45,11 @@ class TestMainWorkflow:
vae=vae,
upscale_by=2.00000004, # Test small float difference doesn't add extra tiles
seed=seed,
steps=10,
steps=5,
cfg=8,
sampler_name="euler",
scheduler="normal",
denoise=0.2,
denoise=0.7,
upscale_model=upscale_model,
mode_type="Chess",
tile_width=512,
@@ -66,11 +63,14 @@ class TestMainWorkflow:
seam_fix_padding=16,
force_uniform_tiles=True,
tiled_decode=False,
batch_size=batch_size,
)
# Save images
im1_filename = image_name_format("upscaled_image1", EXT, batch_size)
im2_filename = image_name_format("upscaled_image2", EXT, batch_size)
sample_dir = test_dirs.sample_images
upscaled_img1_path = sample_dir / IMAGE_1
upscaled_img2_path = sample_dir / IMAGE_2
upscaled_img1_path = sample_dir / CATEGORY / im1_filename
upscaled_img2_path = sample_dir / CATEGORY / im2_filename
save_image(upscaled[0], upscaled_img1_path)
save_image(upscaled[1], upscaled_img2_path)
# Load to account for compression
@@ -83,16 +83,15 @@ class TestMainWorkflow:
im1_upscaled = upscaled[0]
im2_upscaled = upscaled[1]
test_im1_upscaled = load_image(test_image_dir / IMAGE_1)
test_im2_upscaled = load_image(test_image_dir / IMAGE_2)
diff1 = img_tensor_mae(blur(im1_upscaled), blur(test_im1_upscaled))
diff2 = img_tensor_mae(blur(im2_upscaled), blur(test_im2_upscaled))
test_im1 = load_image(test_image_dir / CATEGORY / im1_filename)
test_im2 = load_image(test_image_dir / CATEGORY / im2_filename)
diff1 = img_tensor_mae(blur(im1_upscaled), blur(test_im1))
diff2 = img_tensor_mae(blur(im2_upscaled), blur(test_im2))
# This tolerance is enough to handle both cpu and gpu as the device, as well as jpg compression differences.
logger.info(f"Diff1: {diff1}, Diff2: {diff2}")
assert diff1 < 0.05, "Upscaled Image 1 doesn't match its test image."
assert diff2 < 0.05, "Upscaled Image 2 doesn't match its test image."
assert diff1 < 0.01, "Upscaled Image 1 doesn't match its test image."
assert diff2 < 0.01, "Upscaled Image 2 doesn't match its test image."
def test_upscale_no_upscale(
self,
@@ -100,6 +99,7 @@ class TestMainWorkflow:
loaded_checkpoint,
node_classes,
seed,
batch_size,
test_dirs: DirectoryConfig,
):
"""Generate upscaled images using standard workflow using the no upscale node."""
@@ -121,11 +121,11 @@ class TestMainWorkflow:
negative=negative,
vae=vae,
seed=seed,
steps=10,
steps=5,
cfg=8,
sampler_name="euler",
scheduler="normal",
denoise=0.2,
denoise=0.7,
mode_type="Chess",
tile_width=512,
tile_height=512,
@@ -138,11 +138,14 @@ class TestMainWorkflow:
seam_fix_padding=16,
force_uniform_tiles=True,
tiled_decode=False,
batch_size=batch_size,
)
# Save images
im1_filename = image_name_format("no_upscale_image_1", EXT, batch_size)
im2_filename = image_name_format("no_upscale_image_2", EXT, batch_size)
sample_dir = test_dirs.sample_images
upscaled_img1_path = sample_dir / NO_UPSCALE_IMAGE_1
upscaled_img2_path = sample_dir / NO_UPSCALE_IMAGE_2
upscaled_img1_path = sample_dir / CATEGORY / im1_filename
upscaled_img2_path = sample_dir / CATEGORY / im2_filename
save_image(upscaled[0], upscaled_img1_path)
save_image(upscaled[1], upscaled_img2_path)
# Load to account for compression
@@ -155,15 +158,15 @@ class TestMainWorkflow:
im1_upscaled = upscaled[0]
im2_upscaled = upscaled[1]
test_im1_upscaled = load_image(test_image_dir / NO_UPSCALE_IMAGE_1)
test_im2_upscaled = load_image(test_image_dir / NO_UPSCALE_IMAGE_2)
test_im1 = load_image(test_image_dir / CATEGORY / im1_filename)
test_im2 = load_image(test_image_dir / CATEGORY / im2_filename)
diff1 = img_tensor_mae(blur(im1_upscaled), blur(test_im1_upscaled))
diff2 = img_tensor_mae(blur(im2_upscaled), blur(test_im2_upscaled))
diff1 = img_tensor_mae(blur(im1_upscaled), blur(test_im1))
diff2 = img_tensor_mae(blur(im2_upscaled), blur(test_im2))
# This tolerance is enough to handle both cpu and gpu as the device, as well as jpg compression differences.
logger.info(f"Diff1: {diff1}, Diff2: {diff2}")
assert diff1 < 0.05, "No Upscale Image 1 doesn't match its test image."
assert diff2 < 0.05, "No Upscale Image 2 doesn't match its test image."
assert diff1 < 0.01, f"{im1_filename} doesn't match its test image."
assert diff2 < 0.01, f"{im2_filename} doesn't match its test image."
def test_upscale_with_custom_sampler(
self,
@@ -172,6 +175,7 @@ class TestMainWorkflow:
upscale_model,
node_classes,
seed,
batch_size,
test_dirs: DirectoryConfig,
):
"""Generate upscaled images using standard workflow using the custom sampler node."""
@@ -181,8 +185,8 @@ class TestMainWorkflow:
with torch.inference_mode():
# Setup custom scheduler and sampler
custom_scheduler = node_classes["KarrasScheduler"]
(sigmas,) = execute(custom_scheduler, 20, 14.614642, 0.0291675, 7.0)
(_, sigmas) = execute(node_classes["SplitSigmasDenoise"], sigmas, 0.15)
(sigmas,) = execute(custom_scheduler, 10, 14.614642, 0.0291675, 7.0)
(_, sigmas) = execute(node_classes["SplitSigmasDenoise"], sigmas, 0.7)
custom_sampler = node_classes["KSamplerSelect"]
(sampler,) = execute(custom_sampler, "dpmpp_2m")
@@ -201,7 +205,7 @@ class TestMainWorkflow:
cfg=8,
sampler_name="euler",
scheduler="normal",
denoise=0.2,
denoise=1.0,
upscale_model=upscale_model,
mode_type="Chess",
tile_width=512,
@@ -209,19 +213,22 @@ class TestMainWorkflow:
mask_blur=8,
tile_padding=32,
seam_fix_mode="None",
seam_fix_denoise=1.0,
seam_fix_denoise=0.5,
seam_fix_width=64,
seam_fix_mask_blur=8,
seam_fix_padding=16,
force_uniform_tiles=True,
tiled_decode=False,
batch_size=batch_size,
custom_sampler=sampler,
custom_sigmas=sigmas,
)
# Save images
im1_filename = image_name_format("custom_sampler1", EXT, batch_size)
im2_filename = image_name_format("custom_sampler2", EXT, batch_size)
sample_dir = test_dirs.sample_images
upscaled_img1_path = sample_dir / CUSTOM_SAMPLER_IMAGE_1
upscaled_img2_path = sample_dir / CUSTOM_SAMPLER_IMAGE_2
upscaled_img1_path = sample_dir / CATEGORY / im1_filename
upscaled_img2_path = sample_dir / CATEGORY / im2_filename
save_image(upscaled[0], upscaled_img1_path)
save_image(upscaled[1], upscaled_img2_path)
# Load to account for compression
@@ -233,14 +240,13 @@ class TestMainWorkflow:
test_image_dir = test_dirs.test_images
im1_upscaled = upscaled[0]
im2_upscaled = upscaled[1]
test_im1 = load_image(test_image_dir / CATEGORY / im1_filename)
test_im2 = load_image(test_image_dir / CATEGORY / im2_filename)
test_im1_upscaled = load_image(test_image_dir / CUSTOM_SAMPLER_IMAGE_1)
test_im2_upscaled = load_image(test_image_dir / CUSTOM_SAMPLER_IMAGE_2)
diff1 = img_tensor_mae(blur(im1_upscaled), blur(test_im1_upscaled))
diff2 = img_tensor_mae(blur(im2_upscaled), blur(test_im2_upscaled))
diff1 = img_tensor_mae(blur(im1_upscaled), blur(test_im1))
diff2 = img_tensor_mae(blur(im2_upscaled), blur(test_im2))
# This tolerance is enough to handle both cpu and gpu as the device, as well as jpg compression differences.
logger.info(f"Diff1: {diff1}, Diff2: {diff2}")
assert diff1 < 0.05, "Upscaled Image 1 doesn't match its test image."
assert diff2 < 0.05, "Upscaled Image 2 doesn't match its test image."
assert diff1 < 0.011, f"{im1_filename} doesn't match its test image."
assert diff2 < 0.011, f"{im2_filename} doesn't match its test image."
+49 -37
View File
@@ -6,9 +6,10 @@ import logging
import pathlib
import pytest
import torch
from contextlib import nullcontext
from tensor_utils import img_tensor_mae, blur
from io_utils import save_image, load_image
from io_utils import save_image, load_image, image_name_format
from configs import DirectoryConfig
from fixtures_images import EXT
@@ -16,54 +17,65 @@ from fixtures_images import EXT
CATEGORY = pathlib.Path(pathlib.Path(__file__).stem.removeprefix("test_"))
@pytest.mark.parametrize("batch_size", [1, 2])
def test_minimal_tile_sizes(
base_image, loaded_checkpoint, node_classes, seed, test_dirs: DirectoryConfig
base_image,
loaded_checkpoint,
node_classes,
seed,
batch_size,
test_dirs: DirectoryConfig,
):
"""Test upscaling with minimal tile sizes."""
filename = "non_uniform_tiles"
image, positive, negative = base_image
image = image[0:1] # 1 image for simplicity
model, clip, vae = loaded_checkpoint
with torch.inference_mode():
usdu = node_classes["UltimateSDUpscale"]
(upscaled,) = usdu().upscale(
image=image[0:1],
model=model,
positive=positive,
negative=negative,
vae=vae,
upscale_by=1.5,
seed=seed,
steps=5,
cfg=8,
sampler_name="euler",
scheduler="normal",
denoise=0.15,
upscale_model=None,
mode_type="Chess",
tile_width=512,
tile_height=512,
mask_blur=8,
tile_padding=8,
seam_fix_mode="None",
seam_fix_denoise=1.0,
seam_fix_width=16,
seam_fix_mask_blur=8,
seam_fix_padding=4,
force_uniform_tiles=False,
tiled_decode=False,
)
with pytest.raises(AssertionError) if batch_size > 1 else nullcontext():
usdu = node_classes["UltimateSDUpscale"]
(upscaled,) = usdu().upscale(
image=image,
model=model,
positive=positive,
negative=negative,
vae=vae,
upscale_by=1.5,
seed=seed,
steps=5,
cfg=8,
sampler_name="euler",
scheduler="normal",
denoise=0.6,
upscale_model=None,
mode_type="Chess",
tile_width=512,
tile_height=512,
mask_blur=8,
tile_padding=8,
seam_fix_mode="None",
seam_fix_denoise=1.0,
seam_fix_width=16,
seam_fix_mask_blur=8,
seam_fix_padding=4,
force_uniform_tiles=False, # This should trigger the assertion for batch_size > 1
tiled_decode=False,
batch_size=batch_size,
)
if batch_size > 1:
return # Test passed if assertion was raised
# Save and reload sample image
sample_dir = test_dirs.sample_images
filename_path = CATEGORY / (filename + EXT)
save_image(upscaled[0], sample_dir / filename_path)
upscaled = load_image(sample_dir / filename_path)
filename = CATEGORY / image_name_format("non_uniform_tiles", EXT, batch_size)
save_image(upscaled[0], sample_dir / filename)
upscaled = load_image(sample_dir / filename)
# Compare with reference
test_image_dir = test_dirs.test_images
test_image = load_image(test_image_dir / filename_path)
diff = img_tensor_mae(blur(upscaled), blur(test_image))
test_image = load_image(test_image_dir / filename)
logger = logging.getLogger(__name__)
diff = img_tensor_mae(blur(upscaled), blur(test_image))
logger.info(f"{filename} MAE: {diff}")
assert diff < 0.05, f"{filename} output doesn't match reference"
assert diff < 0.02, f"{filename} does not match reference (MAE {diff})"
+17 -13
View File
@@ -8,7 +8,7 @@ import pytest
import torch
from tensor_utils import img_tensor_mae, blur
from io_utils import save_image, load_image
from io_utils import save_image, load_image, image_name_format
from configs import DirectoryConfig
from fixtures_images import EXT
@@ -16,11 +16,7 @@ from fixtures_images import EXT
CATEGORY = pathlib.Path(pathlib.Path(__file__).stem.removeprefix("test_"))
def image_name_format(prefix: str, mode: str) -> str:
"""Helper for the image name format for the tests below."""
return f"{prefix}_{mode.lower().replace(' ', '_')}{EXT}"
@pytest.mark.parametrize("batch_size", [1, 2])
class TestTilingModes:
def _test_upscale_variant(
self,
@@ -33,6 +29,7 @@ class TestTilingModes:
seam_fix_mode,
seam_fix_denoise,
filename_prefix,
batch_size,
):
"""Helper method to test upscale variants with different parameters."""
logger = logging.getLogger(f"test_{filename_prefix}")
@@ -53,7 +50,7 @@ class TestTilingModes:
cfg=8,
sampler_name="euler",
scheduler="normal",
denoise=0.2,
denoise=0.9,
upscale_model=None,
mode_type=mode_type,
tile_width=512,
@@ -62,11 +59,12 @@ class TestTilingModes:
tile_padding=32,
seam_fix_mode=seam_fix_mode,
seam_fix_denoise=seam_fix_denoise,
seam_fix_width=64,
seam_fix_width=256,
seam_fix_mask_blur=8,
seam_fix_padding=16,
force_uniform_tiles=True,
tiled_decode=False,
batch_size=batch_size,
)
# Save and reload sample image
@@ -79,8 +77,8 @@ class TestTilingModes:
test_image_dir = test_dirs.test_images
test_image = load_image(test_image_dir / filename)
diff = img_tensor_mae(blur(upscaled), blur(test_image))
logger.info(f"{filename_prefix} MAE: {diff}")
assert diff < 0.05, f"{filename_prefix} output doesn't match reference"
logger.info(f"{filename} MAE: {diff}")
assert diff < 0.01, f"{filename} output doesn't match reference"
# "Chess" is tested in the main workflow test
@pytest.mark.parametrize("mode_type", ["Linear", "None"])
@@ -91,10 +89,11 @@ class TestTilingModes:
node_classes,
seed,
mode_type,
batch_size,
test_dirs: DirectoryConfig,
):
"""Test different tiling mode types."""
filename = image_name_format("mode", mode_type)
filename = image_name_format("mode_" + mode_type.lower(), EXT, batch_size)
self._test_upscale_variant(
base_image,
loaded_checkpoint,
@@ -105,6 +104,7 @@ class TestTilingModes:
seam_fix_mode="None",
seam_fix_denoise=1.0,
filename_prefix=filename,
batch_size=batch_size,
)
@pytest.mark.parametrize(
@@ -117,10 +117,13 @@ class TestTilingModes:
node_classes,
seed,
seam_fix_mode,
batch_size,
test_dirs: DirectoryConfig,
):
"""Test different seam fix modes."""
filename = image_name_format("seamfix", seam_fix_mode)
filename = image_name_format(
"seamfix_" + seam_fix_mode.lower().replace(" ", "_"), EXT, batch_size
)
self._test_upscale_variant(
base_image,
loaded_checkpoint,
@@ -129,6 +132,7 @@ class TestTilingModes:
test_dirs,
mode_type="None",
seam_fix_mode=seam_fix_mode,
seam_fix_denoise=0.5,
seam_fix_denoise=0.6,
filename_prefix=filename,
batch_size=batch_size,
)
+21 -10
View File
@@ -56,6 +56,7 @@ def USDU_base_inputs():
# Misc
("force_uniform_tiles", ("BOOLEAN", {"default": True, "tooltip": "Force all tiles to be the same as the set tile size, even when tiles could be smaller. This can help prevent the model from working with irregular tile sizes."})),
("tiled_decode", ("BOOLEAN", {"default": False, "tooltip": "Whether to use tiled decoding when decoding tiles."})),
("batch_size", ("INT", {"default": 1, "min": 1, "max": 4096, "step": 1, "tooltip": "The number of tiles to process in a batch. Higher values can reduce processing time but use more VRAM. Yields different results than individual tiles. Only affects the main redraw step, not the seam fix step."})),
]
optional = []
@@ -107,7 +108,7 @@ class UltimateSDUpscale:
steps, cfg, sampler_name, scheduler, denoise, upscale_model,
mode_type, tile_width, tile_height, mask_blur, tile_padding,
seam_fix_mode, seam_fix_denoise, seam_fix_mask_blur,
seam_fix_width, seam_fix_padding, force_uniform_tiles, tiled_decode,
seam_fix_width, seam_fix_padding, force_uniform_tiles, tiled_decode, batch_size=1,
custom_sampler=None, custom_sigmas=None):
# Store params
self.tile_width = tile_width
@@ -136,13 +137,19 @@ class UltimateSDUpscale:
shared.batch = [tensor_to_pil(image, i) for i in range(len(image))]
shared.batch_as_tensor = image
# Store batch_size for use in processing
self.batch_size = batch_size
print(f"[USDU Batch Debug] UltimateSDUpscale.upscale() using batch_size={batch_size}")
assert batch_size == 1 or force_uniform_tiles, "batch_size greater than 1 requires force_uniform_tiles to be True; all tiles in the batch must be the same size."
# Processing
sdprocessing = StableDiffusionProcessing(
shared.batch[0], model, positive, negative, vae,
seed, steps, cfg, sampler_name, scheduler, denoise, upscale_by, force_uniform_tiles, tiled_decode,
tile_width, tile_height, MODES[self.mode_type], SEAM_FIX_MODES[self.seam_fix_mode],
custom_sampler, custom_sigmas,
custom_sampler, custom_sigmas, batch_size,
)
print(f"[USDU Batch Debug] StableDiffusionProcessing created with batch_size={sdprocessing.batch_size}")
# Disable logging
logger = logging.getLogger()
@@ -188,13 +195,18 @@ class UltimateSDUpscaleNoUpscale(UltimateSDUpscale):
steps, cfg, sampler_name, scheduler, denoise,
mode_type, tile_width, tile_height, mask_blur, tile_padding,
seam_fix_mode, seam_fix_denoise, seam_fix_mask_blur,
seam_fix_width, seam_fix_padding, force_uniform_tiles, tiled_decode):
seam_fix_width, seam_fix_padding, force_uniform_tiles, tiled_decode, batch_size=1):
upscale_by = 1.0
# Store batch_size for use in processing
self.batch_size = batch_size
print(f"[USDU Batch Debug] UltimateSDUpscaleNoUpscale.upscale() received batch_size={batch_size}")
return super().upscale(upscaled_image, model, positive, negative, vae, upscale_by, seed,
steps, cfg, sampler_name, scheduler, denoise, None,
mode_type, tile_width, tile_height, mask_blur, tile_padding,
seam_fix_mode, seam_fix_denoise, seam_fix_mask_blur,
seam_fix_width, seam_fix_padding, force_uniform_tiles, tiled_decode)
seam_fix_width, seam_fix_padding, force_uniform_tiles, tiled_decode, batch_size)
class UltimateSDUpscaleCustomSample(UltimateSDUpscale):
@classmethod
@@ -205,7 +217,7 @@ class UltimateSDUpscaleCustomSample(UltimateSDUpscale):
optional.append(("custom_sampler", ("SAMPLER", {"tooltip": "A custom sampler to use instead of the built-in ComfyUI sampler specified by sampler_name. Only used if both custom_sampler and custom_sigmas are provided."})))
optional.append(("custom_sigmas", ("SIGMAS", {"tooltip": "A custom noise schedule to use during sampling. Only used if both custom_sampler and custom_sigmas are provided."})))
return prepare_inputs(required, optional)
RETURN_TYPES = ("IMAGE",)
FUNCTION = "upscale"
CATEGORY = "image/upscaling"
@@ -216,28 +228,27 @@ class UltimateSDUpscaleCustomSample(UltimateSDUpscale):
steps, cfg, sampler_name, scheduler, denoise,
mode_type, tile_width, tile_height, mask_blur, tile_padding,
seam_fix_mode, seam_fix_denoise, seam_fix_mask_blur,
seam_fix_width, seam_fix_padding, force_uniform_tiles, tiled_decode,
seam_fix_width, seam_fix_padding, force_uniform_tiles, tiled_decode, batch_size=1,
upscale_model=None,
custom_sampler=None, custom_sigmas=None):
return super().upscale(image, model, positive, negative, vae, upscale_by, seed,
steps, cfg, sampler_name, scheduler, denoise, upscale_model,
mode_type, tile_width, tile_height, mask_blur, tile_padding,
seam_fix_mode, seam_fix_denoise, seam_fix_mask_blur,
seam_fix_width, seam_fix_padding, force_uniform_tiles, tiled_decode,
seam_fix_width, seam_fix_padding, force_uniform_tiles, tiled_decode, batch_size,
custom_sampler, custom_sigmas)
# A dictionary that contains all nodes you want to export with their names
# NOTE: names should be globally unique
NODE_CLASS_MAPPINGS = {
"UltimateSDUpscale": UltimateSDUpscale,
"UltimateSDUpscaleNoUpscale": UltimateSDUpscaleNoUpscale,
"UltimateSDUpscaleCustomSample": UltimateSDUpscaleCustomSample
"UltimateSDUpscaleCustomSample": UltimateSDUpscaleCustomSample,
}
# A dictionary that contains the friendly/humanly readable titles for the nodes
NODE_DISPLAY_NAME_MAPPINGS = {
"UltimateSDUpscale": "Ultimate SD Upscale",
"UltimateSDUpscaleNoUpscale": "Ultimate SD Upscale (No Upscale)",
"UltimateSDUpscaleCustomSample": "Ultimate SD Upscale (Custom Sample)"
"UltimateSDUpscaleCustomSample": "Ultimate SD Upscale (Custom Sample)",
}
+396 -42
View File
@@ -1,71 +1,425 @@
# Make some patches to the script
from repositories import ultimate_upscale as usdu
import modules.shared as shared
"""
Refactored USD Upscaler batch processing patch.
Preserves original behavior but:
- Organizes imports and helpers
- Replaces prints with logging
- Factors duplicated logic (tile preparation, batching, decoding)
- Uses functools.wraps when monkey-patching methods
- Adds type hints and docstrings for clarity
"""
from __future__ import annotations
import logging
import math
from PIL import Image
import numpy as np
import torch
from functools import wraps
from typing import Tuple, List, Iterable
from PIL import Image, ImageFilter, ImageDraw
from comfy_extras.nodes_custom_sampler import SamplerCustom
import modules.shared as shared
from nodes import common_ksampler, VAEEncode, VAEDecode, VAEDecodeTiled
from repositories import ultimate_upscale as usdu
import usdu_utils
logger = logging.getLogger(__name__)
logger.addHandler(logging.StreamHandler())
logger.setLevel(logging.INFO)
if (not hasattr(Image, 'Resampling')): # For older versions of Pillow
Image.Resampling = Image
#
# Instead of using multiples of 64, use multiples of 8
#
# Compatibility for older Pillow versions
try:
Image.Resampling # type: ignore
except Exception:
Image.Resampling = Image # type: ignore
def round_length(length, multiple=8):
# -------------------------
# Utility helpers
# -------------------------
def round_length(length: int, multiple: int = 8) -> int:
"""Round length to nearest multiple (default 8)."""
return round(length / multiple) * multiple
# Upscaler
old_init = usdu.USDUpscaler.__init__
def new_init(self, p, image, upscaler_index, save_redraw, save_seams_fix, tile_width, tile_height):
p.width = round_length(image.width * p.upscale_by)
p.height = round_length(image.height * p.upscale_by)
old_init(self, p, image, upscaler_index, save_redraw, save_seams_fix, tile_width, tile_height)
def _sample(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent, denoise, custom_sampler, custom_sigmas):
"""Sampling wrapper that supports a custom sampler or falls back to common_ksampler."""
if custom_sampler is not None and custom_sigmas is not None:
kwargs = dict(
model=model,
add_noise=True,
noise_seed=seed,
cfg=cfg,
positive=positive,
negative=negative,
sampler=custom_sampler,
sigmas=custom_sigmas,
latent_image=latent
)
if hasattr(SamplerCustom, "execute"):
(samples, _) = SamplerCustom.execute(**kwargs)
else:
custom_sample = SamplerCustom()
(samples, _) = getattr(custom_sample, custom_sample.FUNCTION)(**kwargs)
return samples
(samples,) = common_ksampler(model, seed, steps, cfg, sampler_name,
scheduler, positive, negative, latent, denoise=denoise)
return samples
usdu.USDUpscaler.__init__ = new_init
# -------------------------
# Monkey patches for USDUpscaler sizing / redraw / seams fix
# -------------------------
def patch_usdu_upscaler_init():
"""Patch USDUpscaler.__init__ to round upscaler p.width/p.height to multiples."""
old_init = usdu.USDUpscaler.__init__
# Redraw
old_setup_redraw = usdu.USDURedraw.init_draw
@wraps(old_init)
def new_init(self, p, image, upscaler_index, save_redraw, save_seams_fix, tile_width, tile_height):
p.width = round_length(image.width * p.upscale_by)
p.height = round_length(image.height * p.upscale_by)
return old_init(self, p, image, upscaler_index, save_redraw, save_seams_fix, tile_width, tile_height)
usdu.USDUpscaler.__init__ = new_init
def new_setup_redraw(self, p, width, height):
mask, draw = old_setup_redraw(self, p, width, height)
p.width = round_length(self.tile_width + self.padding)
p.height = round_length(self.tile_height + self.padding)
return mask, draw
def patch_usdu_redraw_init():
"""Patch USDURedraw.init_draw to round tile size used for redraw."""
old_init_draw = usdu.USDURedraw.init_draw
@wraps(old_init_draw)
def new_init_draw(self, p, width, height):
mask, draw = old_init_draw(self, p, width, height)
p.width = round_length(self.tile_width + self.padding)
p.height = round_length(self.tile_height + self.padding)
return mask, draw
usdu.USDURedraw.init_draw = new_init_draw
usdu.USDURedraw.init_draw = new_setup_redraw
def patch_usdu_seams_fix_init():
old_init = usdu.USDUSeamsFix.init_draw
# Seams fix
old_setup_seams_fix = usdu.USDUSeamsFix.init_draw
@wraps(old_init)
def new_init(self, p):
old_init(self, p)
p.width = round_length(self.tile_width + self.padding)
p.height = round_length(self.tile_height + self.padding)
usdu.USDUSeamsFix.init_draw = new_init
def new_setup_seams_fix(self, p):
old_setup_seams_fix(self, p)
p.width = round_length(self.tile_width + self.padding)
p.height = round_length(self.tile_height + self.padding)
def patch_usdu_upscale_method():
"""Patch USDUpscaler.upscale to keep shared.batch resized to p.width/p.height."""
old_upscale = usdu.USDUpscaler.upscale
@wraps(old_upscale)
def new_upscale(self):
old_upscale(self)
# Keep shared.batch consistent with the upscaling width/height for subsequent processing.
shared.batch = [self.image] + [
img.resize((self.p.width, self.p.height), resample=Image.LANCZOS)
for img in shared.batch[1:]
]
usdu.USDUpscaler.upscale = new_upscale
usdu.USDUSeamsFix.init_draw = new_setup_seams_fix
# Apply patches
patch_usdu_upscaler_init()
patch_usdu_redraw_init()
patch_usdu_seams_fix_init()
patch_usdu_upscale_method()
#
# Make the script upscale on a batch of images instead of one image
#
# -------------------------
# Patched script.run replacement
# -------------------------
def patched_script_run(self, p, _, tile_width, tile_height, mask_blur, padding, seams_fix_width, seams_fix_denoise, seams_fix_padding,
upscaler_index, save_upscaled_image, redraw_mode, save_seams_fix_image, seams_fix_mask_blur,
seams_fix_type, target_size_type, custom_width, custom_height, custom_scale):
"""
Replacement for usdu.Script.run that preserves the original batch_size
and delegates to the (patched) USDUpscaler and redraw pipeline.
"""
preserved_batch_size = getattr(p, 'batch_size', 1)
logger.info("[USDU Batch Debug] Patched script.run() preserving batch_size=%s", preserved_batch_size)
old_upscale = usdu.USDUpscaler.upscale
# Init (matching original code)
usdu.processing.fix_seed(p)
usdu.devices.torch_gc()
# Keep original file-saving flags as in original code
p.do_not_save_grid = True
p.do_not_save_samples = True
p.inpaint_full_res = False
p.inpainting_fill = 1
p.n_iter = 1
p.batch_size = preserved_batch_size
seed = p.seed
# Init image
init_img = p.init_images[0]
if init_img is None:
return usdu.processing.Processed(p, [], seed, "Empty image")
init_img = usdu.images.flatten(init_img, usdu.shared.opts.img2img_background_color)
# Override size by user choice
if target_size_type == 1:
p.width = custom_width
p.height = custom_height
elif target_size_type == 2:
p.width = math.ceil((init_img.width * custom_scale) / 64) * 64
p.height = math.ceil((init_img.height * custom_scale) / 64) * 64
# Create and run upscaler
upscaler = usdu.USDUpscaler(p, init_img, upscaler_index, save_upscaled_image, save_seams_fix_image, tile_width, tile_height)
upscaler.upscale()
# Drawing & seams fix setup
upscaler.setup_redraw(redraw_mode, padding, mask_blur)
upscaler.setup_seams_fix(seams_fix_padding, seams_fix_denoise, seams_fix_mask_blur, seams_fix_width, seams_fix_type)
upscaler.print_info()
upscaler.add_extra_info()
upscaler.process()
result_images = upscaler.result_images
logger.info("[USDU Batch Debug] Patched script.run() complete, batch_size=%s", p.batch_size)
return usdu.processing.Processed(p, result_images, seed, upscaler.initial_info or "")
def new_upscale(self):
old_upscale(self)
shared.batch = [self.image] + \
[img.resize((self.p.width, self.p.height), resample=Image.LANCZOS) for img in shared.batch[1:]]
# Replace the original script.run with patched version
usdu.Script.run = patched_script_run
usdu.USDUpscaler.upscale = new_upscale
# -------------------------
# Batch processing helpers shared between linear and chess modes
# -------------------------
def _prepare_tile_for_batch(calc_rectangle_fn, current_image: Image.Image, tx: int, ty: int, p) -> Tuple[Image.Image, Tuple[int, int, int, int], Image.Image, Tuple[int, int]]:
"""
Prepare cropped/resized tile, mask, crop-region and tile-size for encoding.
Returns: (cropped_tile, initial_tile_size, tile_mask, tile_size)
"""
tile_mask = Image.new("L", (current_image.width, current_image.height), "black")
tile_draw = ImageDraw.Draw(tile_mask)
tile_draw.rectangle(calc_rectangle_fn(tx, ty), fill="white")
crop_region = usdu_utils.get_crop_region(tile_mask, p.inpaint_full_res_padding)
if p.uniform_tile_mode:
x1, y1, x2, y2 = crop_region
crop_w = x2 - x1
crop_h = y2 - y1
crop_ratio = crop_w / crop_h if crop_h != 0 else 1.0
p_ratio = p.width / p.height if p.height != 0 else 1.0
if crop_ratio > p_ratio:
target_w = crop_w
target_h = round(crop_w / p_ratio)
else:
target_w = round(crop_h * p_ratio)
target_h = crop_h
crop_region, _ = usdu_utils.expand_crop(crop_region, tile_mask.width, tile_mask.height, target_w, target_h)
tile_size = (p.width, p.height)
else:
x1, y1, x2, y2 = crop_region
crop_w = x2 - x1
crop_h = y2 - y1
target_w = math.ceil(crop_w / 8) * 8
target_h = math.ceil(crop_h / 8) * 8
crop_region, tile_size = usdu_utils.expand_crop(crop_region, tile_mask.width, tile_mask.height, target_w, target_h)
# Optional blur
if getattr(p, "mask_blur", 0) > 0:
tile_mask = tile_mask.filter(ImageFilter.GaussianBlur(p.mask_blur))
cropped_tile = current_image.crop(crop_region)
initial_tile_size = cropped_tile.size
if cropped_tile.size != tile_size:
cropped_tile = cropped_tile.resize(tile_size, Image.Resampling.LANCZOS)
return cropped_tile, initial_tile_size, tile_mask, crop_region, tile_size
def _process_batch_tiles(p,
tiles_coords: List[Tuple[int, int]],
images: List[Image.Image],
calc_rectangle_fn,
vae_encoder: VAEEncode,
vae_decoder: VAEDecode,
vae_decoder_tiled: VAEDecodeTiled) -> List[Image.Image]:
"""Encode, sample and decode a batch of tiles and composite them into the given images."""
if not tiles_coords or not images:
return images
batch_tiles = []
batch_masks = []
batch_crop_regions = []
batch_tile_sizes = []
for image in images:
for tx, ty in tiles_coords:
cropped_tile, initial_tile_size, tile_mask, crop_region, tile_size = _prepare_tile_for_batch(calc_rectangle_fn, image, tx, ty, p)
batch_tiles.append((cropped_tile, initial_tile_size))
batch_masks.append(tile_mask)
batch_crop_regions.append(crop_region)
batch_tile_sizes.append(tile_size)
# Encode tiles -> latent
batched_tensors = torch.cat([usdu_utils.pil_to_tensor(tile) for tile, _ in batch_tiles], dim=0)
(latent,) = vae_encoder.encode(p.vae, batched_tensors)
# Condition from first tile (assume same)
first_tile_size = batch_tile_sizes[0]
positive_cropped = usdu_utils.crop_cond(p.positive, batch_crop_regions, p.init_size, images[0].size, first_tile_size)
negative_cropped = usdu_utils.crop_cond(p.negative, batch_crop_regions, p.init_size, images[0].size, first_tile_size)
# Sampling
samples = _sample(p.model, p.seed, p.steps, p.cfg, p.sampler_name, p.scheduler,
positive_cropped, negative_cropped, latent, p.denoise,
p.custom_sampler, p.custom_sigmas)
# Update progress bar if present
if getattr(p, "progress_bar_enabled", False) and getattr(p, "pbar", None) is not None:
p.pbar.update(len(list(tiles_coords)))
# Decode
if not getattr(p, "tiled_decode", False):
(decoded,) = vae_decoder.decode(p.vae, samples)
else:
(decoded,) = vae_decoder_tiled.decode(p.vae, samples, 512)
# Composite tiles back
result_imgs = images
for i, result_img in enumerate(result_imgs):
for j, (tx, ty) in enumerate(tiles_coords):
idx = i * len(tiles_coords) + j
tile_sampled = usdu_utils.tensor_to_pil(decoded, idx)
initial_tile_size = batch_tiles[idx][1]
crop_region = batch_crop_regions[idx]
tile_mask = batch_masks[idx]
if tile_sampled.size != initial_tile_size:
tile_sampled = tile_sampled.resize(initial_tile_size, Image.Resampling.LANCZOS)
image_tile_only = Image.new('RGBA', result_img.size)
image_tile_only.paste(tile_sampled, crop_region[:2])
# Add mask as alpha and composite
temp = image_tile_only.copy()
temp.putalpha(tile_mask)
image_tile_only.paste(temp, image_tile_only)
result = result_img.convert('RGBA')
result.alpha_composite(image_tile_only)
result_img = result.convert('RGB')
result_imgs[i] = result_img
return result_imgs
# -------------------------
# Replace USDURedraw.linear_process and chess_process with batched variants
# -------------------------
def patch_usdu_linear_and_chess_process():
old_linear = usdu.USDURedraw.linear_process
old_chess = usdu.USDURedraw.chess_process
@wraps(old_linear)
def new_linear_process(self, p, image, rows, cols):
batch_size = getattr(p, 'batch_size', 1)
logger.info("[USDU Batch Debug] linear_process called batch_size=%s rows=%s cols=%s total_tiles=%s", batch_size, rows, cols, rows * cols)
if batch_size <= 1:
logger.info("[USDU Batch Debug] Using original single-tile processing (batch_size=%s)", batch_size)
return old_linear(self, p, image, rows, cols)
# Batch mode
vae_encoder = VAEEncode()
vae_decoder = VAEDecode()
vae_decoder_tiled = VAEDecodeTiled()
mask_template, draw_template = self.init_draw(p, image.width, image.height)
tiles_to_process: List[Tuple[int, int]] = []
batch_count = 0
for yi in range(rows):
for xi in range(cols):
if shared.state.interrupted:
break
tiles_to_process.append((xi, yi))
if len(tiles_to_process) >= batch_size or (yi == rows - 1 and xi == cols - 1):
batch_count += 1
logger.info("[USDU Batch Debug] Processing batch #%s with %s tiles: %s", batch_count, len(tiles_to_process), tiles_to_process)
shared.batch = _process_batch_tiles(p, tiles_to_process, shared.batch, self.calc_rectangle, vae_encoder, vae_decoder, vae_decoder_tiled)
tiles_to_process = []
logger.info("[USDU Batch Debug] Linear processing complete. Processed %s batches total.", batch_count)
p.width = image.width
p.height = image.height
return image
@wraps(old_chess)
def new_chess_process(self, p, image, rows, cols):
batch_size = getattr(p, 'batch_size', 1)
if batch_size <= 1:
return old_chess(self, p, image, rows, cols)
vae_encoder = VAEEncode()
vae_decoder = VAEDecode()
vae_decoder_tiled = VAEDecodeTiled()
mask_template, draw_template = self.init_draw(p, image.width, image.height)
# Determine tile "white/black" order
tile_colors = []
for yi in range(rows):
row_colors = []
for xi in range(cols):
color = xi % 2 == 0
if yi > 0 and yi % 2 != 0:
color = not color
row_colors.append(color)
tile_colors.append(row_colors)
# Helper to iterate tiles in chess order: white first, then black
def chess_order_iter(white: bool):
for yi in range(rows):
for xi in range(cols):
if tile_colors[yi][xi] == white:
yield (xi, yi)
# Process white tiles then black tiles
for color in (True, False):
tiles_to_process: List[Tuple[int, int]] = []
for tx, ty in chess_order_iter(color):
if shared.state.interrupted:
break
tiles_to_process.append((tx, ty))
if len(tiles_to_process) >= batch_size:
shared.batch = _process_batch_tiles(p, tiles_to_process, shared.batch, self.calc_rectangle, vae_encoder, vae_decoder, vae_decoder_tiled)
tiles_to_process = []
if tiles_to_process:
shared.batch = _process_batch_tiles(p, tiles_to_process, shared.batch, self.calc_rectangle, vae_encoder, vae_decoder, vae_decoder_tiled)
p.width = image.width
p.height = image.height
return image
usdu.USDURedraw.linear_process = new_linear_process
usdu.USDURedraw.chess_process = new_chess_process
patch_usdu_linear_and_chess_process()
logger.info("USDU batch patches applied successfully.")
+67 -15
View File
@@ -295,19 +295,40 @@ def resize_and_pad_tensor(tensor, width, height, fill=False, blur=False):
return result
def crop_controlnet(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad):
def crop_controlnet(cond_dict, regions, init_size, canvas_size, tile_size, w_pad, h_pad):
"""
Crop controlnet hints to the given region and resize them to the tile size.
If there are multiple regions, the hints will be cropped and resized for each region
and concatenated together in the batch dimension.
Supports multiple regions.
:param cond_dict: dict that contains the conditioning.
:param regions: A tuple or list of tuples of the form (x1, y1, x2, y2) denoting the
upper left and the lower right points of the rectangular region.
:param init_size: The original size of the image that the controlnet hints were generated for.
:param canvas_size: The size of the image that the controlnet hints will be resized to before cropping.
:param tile_size: The size to which each cropped hint will be resized.
:param w_pad: The horizontal padding added to each cropped hint.
:param h_pad: The vertical padding added to each cropped hint.
"""
if "control" not in cond_dict:
return
if not isinstance(regions, list):
regions = [regions]
c = cond_dict["control"]
controlnet = c.copy()
cond_dict["control"] = controlnet
while c is not None:
# hint is shape (B, C, H, W)
hint = controlnet.cond_hint_original
resized_crop = resize_region(region, canvas_size, hint.shape[:-3:-1])
hint = crop_tensor(hint.movedim(1, -1), resized_crop).movedim(-1, 1)
hint = resize_tensor(hint, tile_size[::-1])
controlnet.cond_hint_original = hint
tiled_hints = []
for region in regions:
resized_crop = resize_region(region, canvas_size, hint.shape[:-3:-1])
tiled_hint = crop_tensor(hint.movedim(1, -1), resized_crop).movedim(-1, 1)
tiled_hint = resize_tensor(tiled_hint, tile_size[::-1])
tiled_hints.append(tiled_hint)
controlnet.cond_hint_original = torch.cat(tiled_hints, dim=0)
c = c.previous_controlnet
controlnet.set_previous_controlnet(c.copy() if c is not None else None)
controlnet = controlnet.previous_controlnet
@@ -333,9 +354,18 @@ def region_intersection(region1, region2):
return (x1, y1, x2, y2)
def crop_gligen(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad):
def crop_gligen(cond_dict, regions, init_size, canvas_size, tile_size, w_pad, h_pad):
"""
Crop gligen position conditioning to the given region.
Does not support multiple regions.
"""
if "gligen" not in cond_dict:
return
# Only use first region if multiple regions are given
region = regions if isinstance(regions, tuple) else regions[0]
type, model, cond = cond_dict["gligen"]
if type != "position":
from warnings import warn
@@ -379,10 +409,18 @@ def crop_gligen(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_p
cond_dict["gligen"] = (type, model, cropped)
def crop_area(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad):
def crop_area(cond_dict, regions, init_size, canvas_size, tile_size, w_pad, h_pad):
"""
Crop area conditioning to the given region.
Does not support multiple regions.
"""
if "area" not in cond_dict:
return
# Only use first region if multiple regions are given
region = regions if isinstance(regions, tuple) else regions[0]
# Resize the area conditioning to the canvas size and confine it to the tile region
h, w, y, x = cond_dict["area"]
w, h, x, y = 8 * w, 8 * h, 8 * x, 8 * y
@@ -413,9 +451,18 @@ def crop_area(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad
cond_dict["area"] = (h, w, y, x)
def crop_mask(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad):
def crop_mask(cond_dict, regions, init_size, canvas_size, tile_size, w_pad, h_pad):
"""
Crop the mask conditioning to the given region
Does not support multiple regions.
"""
if "mask" not in cond_dict:
return
# Only use first region if multiple regions are given
region = regions if isinstance(regions, tuple) else regions[0]
mask_tensor = cond_dict["mask"] # (B, H, W)
masks = []
for i in range(mask_tensor.shape[0]):
@@ -443,18 +490,23 @@ def crop_mask(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad
cond_dict["mask"] = torch.cat(masks, dim=0) # (B, H, W)
# Added Flux-Kontext Support crop_reference_latents by TBG ETUR
def crop_reference_latents(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad):
def crop_reference_latents(cond_dict, regions, init_size, canvas_size, tile_size, w_pad, h_pad):
"""
1. Resize each latent to `canvas_size` in latent units.
2. Crop the rectangle `region` (pixel coordinates).
3. Down-sample the crop to latent-space `tile_size`.
Expects a list of BCHW tensors under "reference_latents".
Does not support multiple regions.
"""
latents = cond_dict.get("reference_latents")
if not isinstance(latents, list):
return # nothing to do
# Only use first region if multiple regions are given
region = regions if isinstance(regions, tuple) else regions[0]
k = 8 # down-sample factor from pixel space → latent space (SD-type models)
W_can_px, H_can_px = canvas_size
@@ -503,15 +555,15 @@ def crop_reference_latents(cond_dict, region, init_size, canvas_size, tile_size,
def crop_cond(cond, region, init_size, canvas_size, tile_size, w_pad=0, h_pad=0):
def crop_cond(cond, regions, init_size, canvas_size, tile_size, w_pad=0, h_pad=0):
cropped = []
for emb, x in cond:
cond_dict = x.copy()
n = [emb, cond_dict]
crop_controlnet(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad)
crop_gligen(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad)
crop_area(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad)
crop_mask(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad)
crop_reference_latents(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad)
crop_controlnet(cond_dict, regions, init_size, canvas_size, tile_size, w_pad, h_pad)
crop_gligen(cond_dict, regions, init_size, canvas_size, tile_size, w_pad, h_pad)
crop_area(cond_dict, regions, init_size, canvas_size, tile_size, w_pad, h_pad)
crop_mask(cond_dict, regions, init_size, canvas_size, tile_size, w_pad, h_pad)
crop_reference_latents(cond_dict, regions, init_size, canvas_size, tile_size, w_pad, h_pad)
cropped.append(n)
return cropped