From 4713cf824f97277f6d401c25f8ee3abfb298e8bd Mon Sep 17 00:00:00 2001 From: Sayak Paul Date: Fri, 21 Aug 2026 07:58:58 +0000 Subject: [PATCH] tests: fix CUDA OOM in VidTok slicing/tiling tests TestAutoencoderVidTokSlicingTiling overrode test_enable_disable_tiling and test_enable_disable_slicing with copies that ran their three forward passes outside torch.no_grad(). The retained autograd graphs pushed each test to a 10.9 GiB peak on a 14.74 GiB T4, so with any memory left over from earlier files in the same xdist worker the class OOM'd -- taking test_forward_with_norm_groups down with it. The mixin already implements both tests correctly (no_grad, plus torch.allclose for the disable-comparison instead of the `.all() == .all()` idiom the copies used, which compares two booleans and always passes). Drop the overrides and inherit them. Peak memory on a T4, measured per test: before 10937.3 MiB after 1666.5 MiB Co-Authored-By: Claude Opus 5 (1M context) --- .../test_models_autoencoder_vidtok.py | 54 ------------------- 1 file changed, 54 deletions(-) diff --git a/tests/models/autoencoders/test_models_autoencoder_vidtok.py b/tests/models/autoencoders/test_models_autoencoder_vidtok.py index 323d92b40572..16bfea7efa74 100644 --- a/tests/models/autoencoders/test_models_autoencoder_vidtok.py +++ b/tests/models/autoencoders/test_models_autoencoder_vidtok.py @@ -113,60 +113,6 @@ def test_layerwise_casting_training(self): class TestAutoencoderVidTokSlicingTiling(AutoencoderVidTokTesterConfig, NewAutoencoderTesterMixin): """Slicing and tiling tests for AutoencoderVidTok.""" - def test_enable_disable_tiling(self): - init_dict = self.get_init_dict() - inputs_dict = self.get_dummy_inputs() - - torch.manual_seed(0) - model = self.model_class(**init_dict).to(torch_device) - - torch.manual_seed(0) - output_without_tiling = model(**inputs_dict, generator=torch.manual_seed(0))[0] - - torch.manual_seed(0) - model.enable_tiling() - output_with_tiling = model(**inputs_dict, generator=torch.manual_seed(0))[0] - - assert ( - output_without_tiling.detach().cpu().numpy() - output_with_tiling.detach().cpu().numpy() - ).max() < 0.5, "VAE tiling should not affect the inference results" - - torch.manual_seed(0) - model.disable_tiling() - output_without_tiling_2 = model(**inputs_dict, generator=torch.manual_seed(0))[0] - - assert ( - output_without_tiling.detach().cpu().numpy().all() == output_without_tiling_2.detach().cpu().numpy().all() - ), "Without tiling outputs should match with the outputs when tiling is manually disabled." - - def test_enable_disable_slicing(self): - init_dict = self.get_init_dict() - inputs_dict = self.get_dummy_inputs() - - torch.manual_seed(0) - model = self.model_class(**init_dict).to(torch_device) - inputs_dict.update({"return_dict": False}) - - torch.manual_seed(0) - output_without_slicing = model(**inputs_dict, generator=torch.manual_seed(0))[0] - - torch.manual_seed(0) - model.enable_slicing() - output_with_slicing = model(**inputs_dict, generator=torch.manual_seed(0))[0] - - assert ( - output_without_slicing.detach().cpu().numpy() - output_with_slicing.detach().cpu().numpy() - ).max() < 0.5, "VAE slicing should not affect the inference results" - - torch.manual_seed(0) - model.disable_slicing() - output_without_slicing_2 = model(**inputs_dict, generator=torch.manual_seed(0))[0] - - assert ( - output_without_slicing.detach().cpu().numpy().all() - == output_without_slicing_2.detach().cpu().numpy().all() - ), "Without slicing outputs should match when slicing is manually disabled." - def test_forward_with_norm_groups(self): """VidTok uses layernorm instead of groupnorm.""" init_dict = self.get_init_dict()