Skip to content

Commit 110ab59

Browse files
authored
Merge pull request #3086 from bghira/agent/audio-validation-gather
Gather batch-parallel validation payloads
2 parents 33b0326 + 711b351 commit 110ab59

2 files changed

Lines changed: 88 additions & 7 deletions

File tree

simpletuner/helpers/training/validation.py

Lines changed: 9 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,7 @@
1818
import numpy as np
1919
import torch
2020
import wandb
21+
from accelerate.utils import gather_object
2122
from tqdm import tqdm
2223

2324
from simpletuner.helpers.caching.memory import reclaim_memory
@@ -76,7 +77,7 @@
7677
from simpletuner.helpers.utils.checkpoint_manager import CheckpointManager
7778

7879
logger = logging.getLogger("Validation")
79-
from simpletuner.helpers.training.multi_process import gather_across_processes, should_log, split_across_processes
80+
from simpletuner.helpers.training.multi_process import should_log, split_across_processes
8081

8182
if should_log():
8283
logger.setLevel(os.environ.get("SIMPLETUNER_LOG_LEVEL", "INFO"))
@@ -4191,13 +4192,15 @@ def process_prompts(
41914192

41924193
if use_distributed:
41934194
logger.info(f"[Rank {rank}] Gathering {len(local_payloads)} local payloads")
4194-
gathered_payloads = gather_across_processes(local_payloads)
4195+
gathered_payloads = list(gather_object(local_payloads) or [])
4196+
aggregated_payloads = []
4197+
for payload_group in gathered_payloads:
4198+
if isinstance(payload_group, list):
4199+
aggregated_payloads.extend(payload_group)
4200+
else:
4201+
aggregated_payloads.append(payload_group)
41954202
if not self.accelerator.is_main_process:
41964203
return
4197-
logger.info(
4198-
f"[Rank {rank}] Gathered {len(gathered_payloads)} payload groups: {[len(p) for p in gathered_payloads]}"
4199-
)
4200-
aggregated_payloads = [payload for worker_payloads in gathered_payloads for payload in worker_payloads]
42014204
logger.info(f"[Rank {rank}] Total aggregated payloads: {len(aggregated_payloads)}")
42024205

42034206
aggregated_payloads.sort(key=lambda payload: payload["index"])

tests/test_validation_context_parallel.py

Lines changed: 79 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,11 @@
11
import unittest
22
from contextlib import contextmanager
33
from types import SimpleNamespace
4-
from unittest.mock import patch
4+
from unittest.mock import MagicMock, patch
55

6+
import torch
7+
8+
from simpletuner.helpers.models.common import AudioModelFoundation
69
from simpletuner.helpers.training.validation import Validation, _ValidationWorkItem
710

811

@@ -30,6 +33,23 @@ def _work_items(count: int):
3033
]
3134

3235

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+
3353
class ValidationContextParallelTests(unittest.TestCase):
3454
@contextmanager
3555
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):
85105
self.assertEqual(worker_count, 4)
86106
self.assertEqual([item.index for item in local_items], [0])
87107

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+
88166

89167
if __name__ == "__main__":
90168
unittest.main()

0 commit comments

Comments
 (0)