From 6a4594da4c52043d388599efc20893ef8d07b5a0 Mon Sep 17 00:00:00 2001 From: ssitu Date: Wed, 11 Feb 2026 19:02:29 -0500 Subject: [PATCH] refactor: reuse functions in patch for batched tiles --- test/test_main_workflow.py | 8 ++-- test/test_settings.py | 15 +++---- usdu_patch.py | 90 ++++---------------------------------- 3 files changed, 19 insertions(+), 94 deletions(-) diff --git a/test/test_main_workflow.py b/test/test_main_workflow.py index ee183d0..78b2e7b 100644 --- a/test/test_main_workflow.py +++ b/test/test_main_workflow.py @@ -49,7 +49,7 @@ class TestMainWorkflow: cfg=8, sampler_name="euler", scheduler="normal", - denoise=0.6, + denoise=0.7, upscale_model=upscale_model, mode_type="Chess", tile_width=512, @@ -125,7 +125,7 @@ class TestMainWorkflow: cfg=8, sampler_name="euler", scheduler="normal", - denoise=0.6, + denoise=0.7, mode_type="Chess", tile_width=512, tile_height=512, @@ -185,8 +185,8 @@ class TestMainWorkflow: with torch.inference_mode(): # Setup custom scheduler and sampler custom_scheduler = node_classes["KarrasScheduler"] - (sigmas,) = execute(custom_scheduler, 9, 14.614642, 0.0291675, 7.0) - (_, sigmas) = execute(node_classes["SplitSigmasDenoise"], sigmas, 0.6) + (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") diff --git a/test/test_settings.py b/test/test_settings.py index 3141e4a..f97fbd5 100644 --- a/test/test_settings.py +++ b/test/test_settings.py @@ -68,15 +68,14 @@ def test_minimal_tile_sizes( # Save and reload sample image sample_dir = test_dirs.sample_images - filename = image_name_format("non_uniform_tiles", EXT, batch_size) - filename_path = CATEGORY / filename - 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__) - logger.info(f"{filename_path} MAE: {diff}") - assert diff < 0.01, f"{filename_path} output doesn't match reference" + diff = img_tensor_mae(blur(upscaled), blur(test_image)) + logger.info(f"{filename} MAE: {diff}") + assert diff < 0.02, f"{filename} does not match reference (MAE {diff})" diff --git a/usdu_patch.py b/usdu_patch.py index f90a79d..8d9444a 100644 --- a/usdu_patch.py +++ b/usdu_patch.py @@ -24,6 +24,7 @@ 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()) @@ -45,81 +46,6 @@ def round_length(length: int, multiple: int = 8) -> int: return round(length / multiple) * multiple -def _pil_to_tensor(image: Image.Image) -> torch.Tensor: - """Convert PIL image to CHW-like float tensor in [0,1], with batch dim omitted.""" - arr = np.array(image).astype(np.float32) / 255.0 - t = torch.from_numpy(arr) - # Ensure a channel dimension: HxW -> HxWx1, or HxWxC - if t.ndim == 2: - t = t.unsqueeze(-1) - # Move channel last to channel-first if needed by your VAE? You used unsqueeze(0) previously, - # so preserve the previous behavior: add batch dim at dim=0 but do not permute channels. - t = t.unsqueeze(0) - return t - - -def _tensor_to_pil(img_tensor: torch.Tensor, batch_index: int = 0) -> Image.Image: - """Convert tensor (with batch) to PIL image for a specific batch index.""" - safe = torch.nan_to_num(img_tensor[batch_index]) - arr = (255 * safe.cpu().numpy()).astype(np.uint8) - return Image.fromarray(arr) - - -def _fix_crop_region(region: Tuple[int, int, int, int], image_size: Tuple[int, int]) -> Tuple[int, int, int, int]: - """Adjust crop region to remove trailing pixel if not touching border.""" - image_width, image_height = image_size - x1, y1, x2, y2 = region - if x2 < image_width: - x2 -= 1 - if y2 < image_height: - y2 -= 1 - return x1, y1, x2, y2 - - -def _get_crop_region(mask: Image.Image, pad: int = 0) -> Tuple[int, int, int, int]: - """Get the bounding box of the white region in a mask and pad it.""" - coords = mask.getbbox() - if coords is not None: - x1, y1, x2, y2 = coords - else: - # empty bbox => use inverted (no area) - x1, y1, x2, y2 = mask.width, mask.height, 0, 0 - x1 = max(x1 - pad, 0) - y1 = max(y1 - pad, 0) - x2 = min(x2 + pad, mask.width) - y2 = min(y2 + pad, mask.height) - return _fix_crop_region((x1, y1, x2, y2), (mask.width, mask.height)) - - -def _expand_crop(region: Tuple[int, int, int, int], width: int, height: int, target_width: int, target_height: int) -> Tuple[Tuple[int, int, int, int], Tuple[int, int]]: - """Expand a crop region to target size while keeping it inside image.""" - x1, y1, x2, y2 = region - actual_w = x2 - x1 - actual_h = y2 - y1 - - # Expand horizontally - w_diff = target_width - actual_w - x2 = min(x2 + w_diff // 2, width) - w_diff = target_width - (x2 - x1) - x1 = max(x1 - w_diff, 0) - w_diff = target_width - (x2 - x1) - x2 = min(x2 + w_diff, width) - - # Expand vertically - h_diff = target_height - actual_h - y2 = min(y2 + h_diff // 2, height) - h_diff = target_height - (y2 - y1) - y1 = max(y1 - h_diff, 0) - h_diff = target_height - (y2 - y1) - y2 = min(y2 + h_diff, height) - - return (x1, y1, x2, y2), (target_width, target_height) - - -def _crop_cond(cond, region, init_size, canvas_size, tile_size, w_pad: int = 0, h_pad: int = 0): - """Placeholder simplified crop conditioning for batch processing (keeps original behavior).""" - # This intentionally mirrors your simplified version: returns same conditioning. - return cond def _sample(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent, denoise, custom_sampler, custom_sigmas): @@ -287,7 +213,7 @@ def _prepare_tile_for_batch(calc_rectangle_fn, current_image: Image.Image, tx: i tile_draw = ImageDraw.Draw(tile_mask) tile_draw.rectangle(calc_rectangle_fn(tx, ty), fill="white") - crop_region = _get_crop_region(tile_mask, p.inpaint_full_res_padding) + 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 @@ -301,7 +227,7 @@ def _prepare_tile_for_batch(calc_rectangle_fn, current_image: Image.Image, tx: i else: target_w = round(crop_h * p_ratio) target_h = crop_h - crop_region, _ = _expand_crop(crop_region, tile_mask.width, tile_mask.height, target_w, target_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 @@ -309,7 +235,7 @@ def _prepare_tile_for_batch(calc_rectangle_fn, current_image: Image.Image, tx: i crop_h = y2 - y1 target_w = math.ceil(crop_w / 8) * 8 target_h = math.ceil(crop_h / 8) * 8 - crop_region, tile_size = _expand_crop(crop_region, tile_mask.width, tile_mask.height, target_w, target_h) + 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: @@ -347,14 +273,14 @@ def _process_batch_tiles(p, batch_tile_sizes.append(tile_size) # Encode tiles -> latent - batched_tensors = torch.cat([_pil_to_tensor(tile) for tile, _ in batch_tiles], dim=0) + 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_crop_region = batch_crop_regions[0] first_tile_size = batch_tile_sizes[0] - positive_cropped = _crop_cond(p.positive, first_crop_region, p.init_size, current_image.size, first_tile_size) - negative_cropped = _crop_cond(p.negative, first_crop_region, p.init_size, current_image.size, first_tile_size) + positive_cropped = usdu_utils.crop_cond(p.positive, first_crop_region, p.init_size, current_image.size, first_tile_size) + negative_cropped = usdu_utils.crop_cond(p.negative, first_crop_region, p.init_size, current_image.size, first_tile_size) # Sampling samples = _sample(p.model, p.seed, p.steps, p.cfg, p.sampler_name, p.scheduler, @@ -374,7 +300,7 @@ def _process_batch_tiles(p, # Composite tiles back result_img = current_image for idx, (tx, ty) in enumerate(tiles_coords): - tile_sampled = _tensor_to_pil(decoded, idx) + 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]