1313# See the License for the specific language governing permissions and
1414# limitations under the License.
1515
16+ import pytest
1617import torch
1718
1819from diffusers import MMAudioVAE
1920from diffusers .utils .torch_utils import randn_tensor
2021
21- from ...testing_utils import enable_full_determinism , torch_device
22+ from ...testing_utils import enable_full_determinism , require_accelerator , torch_device
2223from ..testing_utils import (
2324 BaseModelTesterConfig ,
2425 MemoryTesterMixin ,
3031enable_full_determinism ()
3132
3233
34+ MEL_DTYPE = pytest .mark .xfail (
35+ reason = "MMAudio's STFT and mel-filter multiplication require float32." ,
36+ raises = RuntimeError ,
37+ strict = True ,
38+ )
39+ LAYERWISE_CASTING_BACKWARD = pytest .mark .xfail (
40+ reason = "MMAudio's gain multiplication saves weights that are cast to float8 before backward." ,
41+ raises = RuntimeError ,
42+ strict = True ,
43+ )
44+ NORMALIZATION_BUFFER_OFFLOAD = pytest .mark .xfail (
45+ reason = "MMAudio encode/decode bypass the offloading hooks for normalization buffers." ,
46+ raises = RuntimeError ,
47+ strict = True ,
48+ )
49+ INCOMPLETE_DEVICE_MAP = pytest .mark .xfail (
50+ reason = "Automatic device maps omit MMAudio's normalization buffers." ,
51+ raises = ValueError ,
52+ strict = True ,
53+ )
54+
55+
3356class MMAudioVAETesterConfig (BaseModelTesterConfig ):
3457 @property
3558 def model_class (self ):
@@ -79,6 +102,16 @@ def output_shape(self) -> tuple[int, ...]:
79102
80103
81104class TestMMAudioVAEModel (MMAudioVAETesterConfig , ModelTesterMixin ):
105+ @MEL_DTYPE
106+ @require_accelerator
107+ @pytest .mark .skipif (
108+ torch_device not in ["cuda" , "xpu" ],
109+ reason = "float16 and bfloat16 can only be use for inference with an accelerator" ,
110+ )
111+ @pytest .mark .parametrize ("dtype" , [torch .float16 , torch .bfloat16 ], ids = ["fp16" , "bf16" ])
112+ def test_from_save_pretrained_dtype_inference (self , tmp_path , dtype ):
113+ super ().test_from_save_pretrained_dtype_inference (tmp_path , dtype )
114+
82115 def test_latent_shape (self ):
83116 model = self .model_class (** self .get_init_dict ()).to (torch_device ).eval ()
84117 with torch .no_grad ():
@@ -89,7 +122,44 @@ def test_latent_shape(self):
89122
90123
91124class TestMMAudioVAEMemory (MMAudioVAETesterConfig , MemoryTesterMixin ):
92- pass
125+ @MEL_DTYPE
126+ def test_layerwise_casting_memory (self ):
127+ super ().test_layerwise_casting_memory ()
128+
129+ @LAYERWISE_CASTING_BACKWARD
130+ def test_layerwise_casting_training (self ):
131+ super ().test_layerwise_casting_training ()
132+
133+ @NORMALIZATION_BUFFER_OFFLOAD
134+ @pytest .mark .parametrize ("record_stream" , [False , True ])
135+ def test_group_offloading (self , base_model_output , record_stream ):
136+ super ().test_group_offloading (base_model_output , record_stream )
137+
138+ @pytest .mark .parametrize ("record_stream" , [False , True ])
139+ @pytest .mark .parametrize (
140+ "offload_type" , ["block_level" , pytest .param ("leaf_level" , marks = NORMALIZATION_BUFFER_OFFLOAD )]
141+ )
142+ def test_group_offloading_with_layerwise_casting (self , record_stream , offload_type ):
143+ super ().test_group_offloading_with_layerwise_casting (record_stream , offload_type )
144+
145+ @pytest .mark .parametrize ("record_stream" , [False , True ])
146+ @pytest .mark .parametrize (
147+ "offload_type" , ["block_level" , pytest .param ("leaf_level" , marks = NORMALIZATION_BUFFER_OFFLOAD )]
148+ )
149+ def test_group_offloading_with_disk (self , tmp_path , record_stream , offload_type ):
150+ super ().test_group_offloading_with_disk (tmp_path , record_stream , offload_type )
151+
152+ @INCOMPLETE_DEVICE_MAP
153+ def test_cpu_offload (self , base_model_output , tmp_path ):
154+ super ().test_cpu_offload (base_model_output , tmp_path )
155+
156+ @INCOMPLETE_DEVICE_MAP
157+ def test_disk_offload_without_safetensors (self , base_model_output , tmp_path ):
158+ super ().test_disk_offload_without_safetensors (base_model_output , tmp_path )
159+
160+ @INCOMPLETE_DEVICE_MAP
161+ def test_disk_offload_with_safetensors (self , base_model_output , tmp_path ):
162+ super ().test_disk_offload_with_safetensors (base_model_output , tmp_path )
93163
94164
95165class TestMMAudioVAETorchCompile (MMAudioVAETesterConfig , TorchCompileTesterMixin ):
0 commit comments