From f7879edf697d5102dbcd56e7add5a0449943cb29 Mon Sep 17 00:00:00 2001 From: Jordan Thompson Date: Tue, 4 Nov 2025 09:59:13 -0800 Subject: [PATCH] Memory Efficiency Updates --- nodes/TiledUpscaleModel.py | 121 ++++++++++++++++++++++++++----------- 1 file changed, 87 insertions(+), 34 deletions(-) diff --git a/nodes/TiledUpscaleModel.py b/nodes/TiledUpscaleModel.py index 356718b..5f5afc3 100644 --- a/nodes/TiledUpscaleModel.py +++ b/nodes/TiledUpscaleModel.py @@ -16,12 +16,11 @@ def dynamic_tiled_upscale_with_custom_feather( target_height=None, target_width=None, resample_method="lanczos", + device=None, ): if samples.ndim != 4: raise ValueError("Expected samples with shape [B, C, H, W].") - batch_size, channels, in_height, in_width = samples.shape - if target_height is None or target_width is None: raise ValueError("target_height and target_width must be provided.") @@ -35,6 +34,13 @@ def dynamic_tiled_upscale_with_custom_feather( tile_step = tile_size - overlap if tile_size > overlap else tile_size + if device is None: + device = model_management.get_torch_device() + + # Keep the full input on CPU and only move tiles to GPU. + samples = samples.to(output_device) + batch_size, channels, in_height, in_width = samples.shape + scale_y_global = float(target_height) / float(in_height) scale_x_global = float(target_width) / float(in_width) @@ -43,8 +49,8 @@ def dynamic_tiled_upscale_with_custom_feather( for batch_index in range(batch_size): source = samples[batch_index : batch_index + 1] - output_for_batch = None - weight_for_batch = None + output_accumulator = None + weight_map = None y_position = 0 while y_position < in_height: @@ -53,24 +59,32 @@ def dynamic_tiled_upscale_with_custom_feather( y_end = min(y_position + tile_size, in_height) x_end = min(x_position + tile_size, in_width) - tile_source = source[:, :, y_position:y_end, x_position:x_end] - tile_output_native = function(tile_source).to(output_device) + # CPU -> GPU per-tile + tile_source_cpu = source[:, :, y_position:y_end, x_position:x_end] + tile_source = tile_source_cpu.to(device, non_blocking=False) - if output_for_batch is None: + tile_output_native = function(tile_source) + + if blended_output is None: out_channels = tile_output_native.shape[1] - output_for_batch = torch.zeros( - (1, out_channels, target_height, target_width), + blended_output = torch.zeros( + (batch_size, out_channels, target_height, target_width), device=output_device, dtype=tile_output_native.dtype, ) - weight_for_batch = torch.zeros_like(output_for_batch) - if blended_output is None: - blended_output = torch.zeros( - (batch_size, out_channels, target_height, target_width), - device=output_device, - dtype=tile_output_native.dtype, - ) + if output_accumulator is None: + output_accumulator = torch.zeros( + (1, tile_output_native.shape[1], target_height, target_width), + device=output_device, + dtype=tile_output_native.dtype, + ) + # Single-channel weight map (broadcast to all channels) + weight_map = torch.zeros( + (1, 1, target_height, target_width), + device=output_device, + dtype=tile_output_native.dtype, + ) out_y_start = int(round(y_position * target_height / in_height)) out_y_end = int(round(y_end * target_height / in_height)) @@ -84,7 +98,7 @@ def dynamic_tiled_upscale_with_custom_feather( tile_output_native.shape[2] != tile_target_height or tile_output_native.shape[3] != tile_target_width ): - tile_output = comfy.utils.common_upscale( + tile_output_gpu = comfy.utils.common_upscale( tile_output_native, tile_target_width, tile_target_height, @@ -92,9 +106,9 @@ def dynamic_tiled_upscale_with_custom_feather( "disabled", ) else: - tile_output = tile_output_native + tile_output_gpu = tile_output_native - mask = torch.ones_like(tile_output) + tile_output = tile_output_gpu.to(output_device) if feather is None or feather <= 0: feather_pixels_y = int(round(overlap * scale_y_global)) @@ -103,6 +117,12 @@ def dynamic_tiled_upscale_with_custom_feather( feather_pixels_y = int(feather) feather_pixels_x = int(feather) + mask = torch.ones( + (1, 1, tile_output.shape[2], tile_output.shape[3]), + device=output_device, + dtype=tile_output.dtype, + ) + if feather_pixels_y > 0: max_vertical = tile_output.shape[2] // 2 feather_pixels_y = min(feather_pixels_y, max_vertical) @@ -130,27 +150,34 @@ def dynamic_tiled_upscale_with_custom_feather( out_y_end = out_y_start + tile_output.shape[2] out_x_end = out_x_start + tile_output.shape[3] - output_for_batch[:, :, out_y_start:out_y_end, out_x_start:out_x_end] += ( + output_accumulator[:, :, out_y_start:out_y_end, out_x_start:out_x_end] += ( tile_output * mask ) - weight_for_batch[:, :, out_y_start:out_y_end, out_x_start:out_x_end] += mask + weight_map[:, :, out_y_start:out_y_end, out_x_start:out_x_end] += mask if pbar is not None: pbar.update(1) + del tile_output_gpu, tile_output_native, tile_source + torch.cuda.empty_cache() + x_position += tile_step y_position += tile_step - weight_for_batch = torch.where( - weight_for_batch == 0.0, - torch.ones_like(weight_for_batch), - weight_for_batch, + safe_weight_map = torch.where( + weight_map == 0.0, + torch.ones_like(weight_map), + weight_map, ) - output_for_batch = output_for_batch / weight_for_batch - blended_output[batch_index : batch_index + 1] = output_for_batch + blended_output[batch_index : batch_index + 1] = ( + output_accumulator / safe_weight_map + ) - return blended_output + del output_accumulator, weight_map + torch.cuda.empty_cache() + + return blended_output.to(output_device) class WASTiledImageUpscaleWithModel: @@ -213,6 +240,13 @@ class WASTiledImageUpscaleWithModel: "tooltip": "Resampling kernel used to reach the final upscale_factor resolution.", }, ), + "clear_comfy_memory": ( + "BOOLEAN", + { + "default": False, + "tooltip": "If enabled, will unload all model, and do a soft cache dump followed by a hard cache dump. Aggressively frees VRAM by unloading other models.", + }, + ), } } @@ -229,9 +263,29 @@ class WASTiledImageUpscaleWithModel: overlap, feather, resample_method, + clear_comfy_memory, ): device = model_management.get_torch_device() + if clear_comfy_memory: + try: + if hasattr(model_management, "unload_all_models"): + model_management.unload_all_models() + print("[WASTiledImageUpscaleWithModel] Cleared all models") + else: + print("[WASTiledImageUpscaleWithModel] unload_all_models not found") + if hasattr(model_management, "soft_empty_cache"): + model_management.soft_empty_cache(True) + print("[WASTiledImageUpscaleWithModel] Cleared soft cache") + else: + print("[WASTiledImageUpscaleWithModel] soft_empty_cache not found") + try: + torch.cuda.empty_cache() + except Exception as exception: + print(f"[WASTiledImageUpscaleWithModel] Failed to empty PyTorch cache: {exception}") + except Exception as exception: + print(f"[WASTiledImageUpscaleWithModel] Failed to unload all comfy models and cache: {exception}") + scale_estimate = getattr(upscale_model, "scale", 4.0) element_size = image.element_size() @@ -254,13 +308,12 @@ class WASTiledImageUpscaleWithModel: target_height = max(1, int(round(in_h * upscale_factor))) target_width = max(1, int(round(in_w * upscale_factor))) - input_image = image.movedim(-1, -3).to(device) + input_image = image.movedim(-1, -3).to("cpu") current_tile_size = int(tile_size) minimum_tile_size = 64 upscale_result = None - output_device = device oom = True last_exception = None @@ -281,12 +334,13 @@ class WASTiledImageUpscaleWithModel: function=lambda a: upscale_model(a), tile_size=current_tile_size, overlap=overlap, - output_device=output_device, + output_device="cpu", pbar=progress, feather=feather, target_height=target_height, target_width=target_width, resample_method=resample_method, + device=device, ) oom = False @@ -299,9 +353,8 @@ class WASTiledImageUpscaleWithModel: upscale_model.to("cpu") - upscale_result = torch.clamp( - upscale_result.movedim(-3, -1), min=0.0, max=1.0 - ) + upscale_result = torch.clamp(upscale_result, min=0.0, max=1.0) + upscale_result = upscale_result.movedim(-3, -1) return (upscale_result,)