"""Tiled per-component SR with measured CUDA costs and bounded concurrent waves.""" import gc from concurrent.futures import ThreadPoolExecutor import numpy as np from PIL import Image from .resources import parallelism, snapshot def upscale_parts(parts, model, log): import torch tile = 256 device = "cuda" def predict(image): array = np.array(image.convert("RGB"), dtype=np.float32) / 255 tensor = torch.from_numpy(array).permute(2, 0, 1).unsqueeze(0).to(device) with torch.inference_mode(): output = model(tensor).clamp(0, 1) return Image.fromarray((output[0].permute(1, 2, 0).cpu().numpy() * 255).round().astype(np.uint8)) # Profile one maximal padded tile before admitting simultaneous requests. while True: try: torch.cuda.empty_cache() free, _ = torch.cuda.mem_get_info() before = torch.cuda.memory_allocated() torch.cuda.reset_peak_memory_stats() predict(Image.new("RGB", (tile + 32, tile + 32))) torch.cuda.synchronize() per_gpu = max(torch.cuda.max_memory_allocated() - before, 64 * 1024**2) per_ram = max(max(p.visual.material.baseColorTexture.width * p.visual.material.baseColorTexture.height * 4 * 16 * 4 for _, p in parts), 128 * 1024**2) parallelism(snapshot(free), per_ram, per_gpu, len(parts)) break except (torch.cuda.OutOfMemoryError, MemoryError): gc.collect() torch.cuda.empty_cache() if tile <= 32: raise MemoryError("A single minimum-size SR tile cannot fit with the 10% reserve.") from None tile //= 2 def worker(item): key, part = item source = part.visual.material.baseColorTexture.convert("RGBA") output = Image.new("RGBA", (source.width * 4, source.height * 4)) stream = torch.cuda.Stream() with torch.cuda.stream(stream): for y in range(0, source.height, tile): for x in range(0, source.width, tile): right, bottom = min(x + tile, source.width), min(y + tile, source.height) left_pad, top_pad = max(0, x - 16), max(0, y - 16) right_pad, bottom_pad = min(source.width, right + 16), min(source.height, bottom + 16) patch = predict(source.crop((left_pad, top_pad, right_pad, bottom_pad))) patch = patch.crop(((x-left_pad)*4, (y-top_pad)*4, (right-left_pad)*4, (bottom-top_pad)*4)) output.paste(patch, (x * 4, y * 4)) stream.synchronize() output.putalpha(source.getchannel("A").resize(output.size, Image.Resampling.BILINEAR)) return key, output remaining = list(parts) completed = {} cap = len(parts) while remaining: torch.cuda.empty_cache() free, _ = torch.cuda.mem_get_info() count = min(cap, parallelism(snapshot(free), per_ram, per_gpu, len(remaining))) log(f"Upscaling {count} component(s) concurrently; tile={tile}, VRAM/task={per_gpu}, reserve=10%") wave, remaining = remaining[:count], remaining[count:] retries = [] with ThreadPoolExecutor(max_workers=count) as pool: futures = [(item, pool.submit(worker, item)) for item in wave] for item, future in futures: try: key, texture = future.result() completed[key] = texture except torch.cuda.OutOfMemoryError as error: error.__traceback__ = None retries.append(item) if retries: gc.collect() torch.cuda.empty_cache() if count == 1: if tile <= 32: raise MemoryError("Texture SR ran out of memory at minimum tile size.") tile //= 2 cap = max(1, count // 2) remaining = retries + remaining for key, part in parts: part.visual.material.baseColorTexture = completed[key] return {"tile_size": tile, "measured_vram_per_task": per_gpu, "ram_budget_per_task": per_ram, "reserve": 0.10, "scale": 4, "model": "RealESRGAN_x4plus", "component_count": len(parts)}