mantrakp's picture
Isolate reference inference in a dedicated ZeroGPU worker
5c331a4 verified
Raw History Blame Contribute Delete
4.38 kB
"""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)}