Merge 6008a10429 into a9ab2b62da
This commit is contained in:
commit
83ed690f7b
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Reference in New Issue