Fix tiled VAE wrapper review regressions

This commit is contained in:
xmarre
2026-04-16 05:35:27 +02:00
parent bd249781b3
commit 5c623b9cdb
+92 -9
View File
@@ -2,10 +2,12 @@ from __future__ import annotations
import contextlib
import functools
import inspect
import json
import logging
import os
import struct
import threading
from typing import Any, Callable
import torch
@@ -39,6 +41,8 @@ _WARNED_PICKLE_GPU_PATHS: set[str] = set()
_SAFE_TENSORS_COMPONENT_CACHE_MAX = 32
_STICKY_PROTECTION_VRAM_FLOOR_RATIO = 0.125
_STICKY_PROTECTION_VRAM_FLOOR_CEIL_BYTES = 16 * 1024 ** 3
_TILED_VAE_MEMORY_LOCK_ATTR = "_gpu_resident_loader_tiled_memory_lock"
_TILED_VAE_LOCK_INIT = threading.Lock()
_SAFETENSORS_DTYPE_MAP = {
"BOOL": torch.bool,
"U8": torch.uint8,
@@ -927,6 +931,14 @@ def _temporary_tiled_vae_memory_estimate(
if not callable(original):
yield
return
had_instance_attr = memory_attr in getattr(self, "__dict__", {})
lock = getattr(self, _TILED_VAE_MEMORY_LOCK_ATTR, None)
if lock is None:
with _TILED_VAE_LOCK_INIT:
lock = getattr(self, _TILED_VAE_MEMORY_LOCK_ATTR, None)
if lock is None:
lock = threading.RLock()
setattr(self, _TILED_VAE_MEMORY_LOCK_ATTR, lock)
def estimated(shape, dtype, *args, **kwargs):
shapes = _tiled_vae_memory_shapes(
@@ -940,18 +952,62 @@ def _temporary_tiled_vae_memory_estimate(
)
return max(int(original(candidate, dtype, *args, **kwargs)) for candidate in shapes)
setattr(self, memory_attr, estimated)
try:
yield
finally:
setattr(self, memory_attr, original)
with lock:
setattr(self, memory_attr, estimated)
try:
yield
finally:
if had_instance_attr:
setattr(self, memory_attr, original)
else:
delattr(self, memory_attr)
@functools.lru_cache(maxsize=None)
def _tiled_vae_supported_kwargs(func: Callable[..., Any]) -> frozenset[str]:
return frozenset(inspect.signature(func).parameters)
def _call_tiled_vae(
func: Callable[..., Any],
self,
data,
*,
tile_x=None,
tile_y=None,
overlap=None,
tile_t=None,
overlap_t=None,
):
kwargs = {}
supported_kwargs = _tiled_vae_supported_kwargs(func)
if "tile_x" in supported_kwargs:
kwargs["tile_x"] = tile_x
if "tile_y" in supported_kwargs:
kwargs["tile_y"] = tile_y
if "overlap" in supported_kwargs:
kwargs["overlap"] = overlap
if "tile_t" in supported_kwargs:
kwargs["tile_t"] = tile_t
if "overlap_t" in supported_kwargs:
kwargs["overlap_t"] = overlap_t
return func(self, data, **kwargs)
def _wrap_vae_encode_tiled(func: Callable[..., Any]) -> Callable[..., Any]:
@functools.wraps(func)
def wrapper(self, pixel_samples, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None):
if REGISTRY.get_policy() != "sticky_gpu":
return func(self, pixel_samples, tile_x=tile_x, tile_y=tile_y, overlap=overlap, tile_t=tile_t, overlap_t=overlap_t)
return _call_tiled_vae(
func,
self,
pixel_samples,
tile_x=tile_x,
tile_y=tile_y,
overlap=overlap,
tile_t=tile_t,
overlap_t=overlap_t,
)
with _temporary_tiled_vae_memory_estimate(
self,
@@ -960,7 +1016,16 @@ def _wrap_vae_encode_tiled(func: Callable[..., Any]) -> Callable[..., Any]:
tile_y=tile_y,
tile_t=tile_t,
):
return func(self, pixel_samples, tile_x=tile_x, tile_y=tile_y, overlap=overlap, tile_t=tile_t, overlap_t=overlap_t)
return _call_tiled_vae(
func,
self,
pixel_samples,
tile_x=tile_x,
tile_y=tile_y,
overlap=overlap,
tile_t=tile_t,
overlap_t=overlap_t,
)
return wrapper
@@ -969,7 +1034,16 @@ def _wrap_vae_decode_tiled(func: Callable[..., Any]) -> Callable[..., Any]:
@functools.wraps(func)
def wrapper(self, samples, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None):
if REGISTRY.get_policy() != "sticky_gpu":
return func(self, samples, tile_x=tile_x, tile_y=tile_y, overlap=overlap, tile_t=tile_t, overlap_t=overlap_t)
return _call_tiled_vae(
func,
self,
samples,
tile_x=tile_x,
tile_y=tile_y,
overlap=overlap,
tile_t=tile_t,
overlap_t=overlap_t,
)
with _temporary_tiled_vae_memory_estimate(
self,
@@ -978,7 +1052,16 @@ def _wrap_vae_decode_tiled(func: Callable[..., Any]) -> Callable[..., Any]:
tile_y=tile_y,
tile_t=tile_t,
):
return func(self, samples, tile_x=tile_x, tile_y=tile_y, overlap=overlap, tile_t=tile_t, overlap_t=overlap_t)
return _call_tiled_vae(
func,
self,
samples,
tile_x=tile_x,
tile_y=tile_y,
overlap=overlap,
tile_t=tile_t,
overlap_t=overlap_t,
)
return wrapper