This commit is contained in:
chelsealong 2026-08-15 17:31:47 +00:00 committed by GitHub
commit 83ed690f7b
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 66 additions and 0 deletions

View File

@ -1250,6 +1250,9 @@ class VAE:
if self.handles_tiling:
tile = 256 // self.spacial_compression_decode()
overlap = tile // 4
# decode_tiled is identical to decode for these VAEs, so freeing the
# memory other models hold onto is the only thing that can make the retry succeed.
model_management.free_memory(1e30, self.device, keep_loaded=[model_management.LoadedModel(self.patcher)])
pixel_samples = self._decode_tiled_owned(samples_in, tile_x=tile, tile_y=tile, overlap=overlap)
else:
pixel_samples = self.decode_tiled_(samples_in)
@ -1257,6 +1260,9 @@ class VAE:
tile = 256 // self.spacial_compression_decode()
overlap = tile // 4
if self.handles_tiling:
# decode_tiled is identical to decode for these VAEs, so freeing the
# memory other models hold onto is the only thing that can make the retry succeed.
model_management.free_memory(1e30, self.device, keep_loaded=[model_management.LoadedModel(self.patcher)])
memory_used = self.memory_used_decode(self._tile_bounded_shape(samples_in.shape, tile, tile, None), self.vae_dtype)
model_management.load_models_gpu([self.patcher], memory_required=memory_used, force_full_load=self.disable_offload)
pixel_samples = self._decode_tiled_owned(samples_in, tile_x=tile, tile_y=tile, overlap=overlap)
@ -1374,6 +1380,9 @@ class VAE:
tile = 256
overlap = tile // 4
if self.handles_tiling:
# encode_tiled is identical to encode for these VAEs, so freeing the
# memory other models hold onto is the only thing that can make the retry succeed.
model_management.free_memory(1e30, self.device, keep_loaded=[model_management.LoadedModel(self.patcher)])
samples = self._encode_tiled_owned(pixel_samples, tile_x=tile, tile_y=tile, overlap=overlap)
else:
samples = self.encode_tiled_3d(pixel_samples, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap))

View File

@ -392,6 +392,63 @@ def test_seedvr2_3d_routes_to_owned_encode_tiled_on_oom():
)
def test_handles_tiling_decode_oom_frees_memory_before_retry():
wrapper = seedvr_vae_mod.VideoAutoencoderKLWrapper.__new__(
seedvr_vae_mod.VideoAutoencoderKLWrapper)
vae = _make_vae(wrapper, latent_channels=_LATENT_CHANNELS, latent_dim=3)
call_order = []
free_memory_call = MagicMock(side_effect=lambda *a, **k: call_order.append("free_memory"))
seedvr2_call = MagicMock(
side_effect=lambda *a, **k: call_order.append("decode_tiled_owned") or torch.zeros(1, 3, 9, 64, 64))
mm = sd_mod.model_management
with ExitStack() as stack:
stack.enter_context(patch.object(mm, "raise_non_oom", lambda e: None))
stack.enter_context(patch.object(mm, "load_models_gpu", lambda *a, **k: None))
stack.enter_context(patch.object(mm, "soft_empty_cache", lambda: None))
stack.enter_context(patch.object(mm, "free_memory", free_memory_call))
stack.enter_context(patch.object(sd_mod.VAE, "_decode_tiled_owned", seedvr2_call))
stack.enter_context(patch.object(
seedvr_vae_mod.VideoAutoencoderKLWrapper, "decode",
side_effect=_force_oom))
vae.decode(torch.zeros(1, _LATENT_CHANNELS * 3, 8, 8))
assert call_order == ["free_memory", "decode_tiled_owned"], (
"decode_tiled retry for a handles_tiling VAE must free other models' "
f"memory before retrying; call order was {call_order}"
)
assert free_memory_call.call_args.kwargs["keep_loaded"][0].model is vae.patcher
def test_handles_tiling_encode_oom_frees_memory_before_retry():
vae = _make_seedvr2_vae_fallback()
pixel_samples = torch.zeros((1, 8, 64, 64, 3))
call_order = []
free_memory_call = MagicMock(side_effect=lambda *a, **k: call_order.append("free_memory"))
seedvr2_call = MagicMock(
side_effect=lambda *a, **k: call_order.append("encode_tiled_owned") or torch.zeros(1, _LATENT_CHANNELS, 2, 8, 8))
mm = sd_mod.model_management
with ExitStack() as stack:
stack.enter_context(patch.object(mm, "raise_non_oom", lambda e: None))
stack.enter_context(patch.object(mm, "load_models_gpu", lambda *a, **k: None))
stack.enter_context(patch.object(mm, "soft_empty_cache", lambda: None))
stack.enter_context(patch.object(mm, "free_memory", free_memory_call))
stack.enter_context(patch.object(sd_mod.VAE, "_encode_tiled_owned", seedvr2_call))
stack.enter_context(patch.object(
seedvr_vae_mod.VideoAutoencoderKLWrapper, "encode",
side_effect=_force_regular_encode_oom))
vae.encode(pixel_samples)
assert call_order == ["free_memory", "encode_tiled_owned"], (
"encode_tiled retry for a handles_tiling VAE must free other models' "
f"memory before retrying; call order was {call_order}"
)
assert free_memory_call.call_args.kwargs["keep_loaded"][0].model is vae.patcher
def test_non_seedvr2_encode_tiled_3d_default_overlap_is_concrete():
vae = _make_non_seedvr2_vae_fallback()
vae.downscale_ratio = (lambda a: max(1, a // 4), 8, 8)