|
1 | 1 | import unittest |
2 | 2 | from contextlib import contextmanager |
3 | 3 | from types import SimpleNamespace |
4 | | -from unittest.mock import patch |
| 4 | +from unittest.mock import MagicMock, patch |
5 | 5 |
|
| 6 | +import torch |
| 7 | + |
| 8 | +from simpletuner.helpers.models.common import AudioModelFoundation |
6 | 9 | from simpletuner.helpers.training.validation import Validation, _ValidationWorkItem |
7 | 10 |
|
8 | 11 |
|
@@ -30,6 +33,23 @@ def _work_items(count: int): |
30 | 33 | ] |
31 | 34 |
|
32 | 35 |
|
| 36 | +class DummyAudioModel(AudioModelFoundation): |
| 37 | + def validation_audio_sample_rate(self): |
| 38 | + return 44100 |
| 39 | + |
| 40 | + def _encode_prompts(self, prompts, is_negative_prompt=False): |
| 41 | + return {} |
| 42 | + |
| 43 | + def convert_text_embed_for_pipeline(self, text_embedding): |
| 44 | + return {} |
| 45 | + |
| 46 | + def convert_negative_text_embed_for_pipeline(self, text_embedding): |
| 47 | + return {} |
| 48 | + |
| 49 | + def model_predict(self, prepared_batch): |
| 50 | + return None |
| 51 | + |
| 52 | + |
33 | 53 | class ValidationContextParallelTests(unittest.TestCase): |
34 | 54 | @contextmanager |
35 | 55 | def _patch_cp(self, *, data_rank: int, data_local_rank: int, data_parallel_size: int = 4): |
@@ -85,6 +105,64 @@ def test_context_parallel_keeps_single_prompt_distributed(self): |
85 | 105 | self.assertEqual(worker_count, 4) |
86 | 106 | self.assertEqual([item.index for item in local_items], [0]) |
87 | 107 |
|
| 108 | + def test_batch_parallel_gathers_audio_payloads_from_peer_ranks(self): |
| 109 | + validation = Validation.__new__(Validation) |
| 110 | + validation.accelerator = SimpleNamespace( |
| 111 | + num_processes=2, |
| 112 | + process_index=0, |
| 113 | + is_main_process=True, |
| 114 | + ) |
| 115 | + validation.config = SimpleNamespace(validation_multigpu="batch-parallel") |
| 116 | + validation.model = DummyAudioModel.__new__(DummyAudioModel) |
| 117 | + validation.validation_prompt_metadata = { |
| 118 | + "validation_prompts": ["prompt 0", "prompt 1"], |
| 119 | + "validation_shortnames": ["song_0", "song_1"], |
| 120 | + } |
| 121 | + validation.validation_image_inputs = None |
| 122 | + validation.validation_prompt_dict = None |
| 123 | + validation.validation_resolutions = [(0, 0)] |
| 124 | + validation.save_dir = "validation_images" |
| 125 | + validation.eval_scores = {} |
| 126 | + validation.validation_video_paths = {} |
| 127 | + validation.evaluation_result = None |
| 128 | + validation._check_abort = MagicMock() |
| 129 | + validation._use_context_parallel_validation = MagicMock(return_value=False) |
| 130 | + validation._split_validation_work_items = MagicMock(return_value=([_work_items(2)[0]], True, 2)) |
| 131 | + validation._should_publish_validation_payloads = MagicMock(return_value=True) |
| 132 | + |
| 133 | + rank0_payload = { |
| 134 | + "index": 0, |
| 135 | + "shortname": "song_0", |
| 136 | + "decorated_shortname": "song_0", |
| 137 | + "prompt": "prompt 0", |
| 138 | + "stitched": [], |
| 139 | + "checkpoint": [], |
| 140 | + "audio": Validation._serialise_media_list([torch.zeros(1, 4)]), |
| 141 | + } |
| 142 | + rank1_payload = { |
| 143 | + "index": 1, |
| 144 | + "shortname": "song_1", |
| 145 | + "decorated_shortname": "song_1", |
| 146 | + "prompt": "prompt 1", |
| 147 | + "stitched": [], |
| 148 | + "checkpoint": [], |
| 149 | + "audio": Validation._serialise_media_list([torch.ones(1, 4)]), |
| 150 | + } |
| 151 | + validation._execute_validation_work_item = MagicMock(return_value=rank0_payload) |
| 152 | + |
| 153 | + with ( |
| 154 | + patch("simpletuner.helpers.training.validation.gather_object", return_value=[[rank0_payload], [rank1_payload]]), |
| 155 | + patch("simpletuner.helpers.training.validation.validation_audio.save_audio") as save_audio, |
| 156 | + patch("simpletuner.helpers.training.validation.validation_audio.log_audio_to_webhook"), |
| 157 | + patch("simpletuner.helpers.training.validation.validation_audio.log_audio_to_trackers"), |
| 158 | + ): |
| 159 | + validation.process_prompts(validation_type="intermediary") |
| 160 | + |
| 161 | + self.assertEqual(sorted(validation.validation_audios.keys()), ["song_0", "song_1"]) |
| 162 | + torch.testing.assert_close(validation.validation_audios["song_0"][0], torch.zeros(1, 4)) |
| 163 | + torch.testing.assert_close(validation.validation_audios["song_1"][0], torch.ones(1, 4)) |
| 164 | + self.assertEqual(save_audio.call_count, 2) |
| 165 | + |
88 | 166 |
|
89 | 167 | if __name__ == "__main__": |
90 | 168 | unittest.main() |
0 commit comments