RustyMark/dots.tts
0
1from __future__ import annotations2 3import math4import multiprocessing as mp5from collections.abc import Iterable6from copy import deepcopy7 8from torch.utils.data import DataLoader, IterableDataset, get_worker_info9 10from dots_tts.data.batchers import OnlineBatcher11from dots_tts.utils.profiling import ensure_data_profiler12from dots_tts.data.source_adapters.base_adapter import BaseSourceAdapter, SourceContext13 14_TRACKING_KEY = "__tracking_state__"15_RESUME_TOPOLOGY_KEY = "resume_topology"16 17 18def identity_collate(sample):19 return sample20 21 22class StreamingSampleDataset(IterableDataset):23 def __init__(24 self,25 *,26 source: BaseSourceAdapter,27 rank: int,28 world_size: int,29 seed: int,30 ):31 self.source = source32 self.rank = int(rank)33 self.world_size = int(world_size)34 self.seed = int(seed)35 self._epoch = mp.Value("q", 0)36 self._pending_resume_state: dict | None = None37 38 def load_state_dict(self, state: dict | None) -> None:39 self._pending_resume_state = deepcopy(state) if state else None40 41 def set_epoch(self, epoch: int) -> None:42 with self._epoch.get_lock():43 self._epoch.value = int(epoch)44 45 def _current_epoch(self) -> int:46 with self._epoch.get_lock():47 return int(self._epoch.value)48 49 def _take_resume_state(self, epoch: int) -> dict | None:50 if (51 self._pending_resume_state is None52 or int(self._pending_resume_state.get("epoch", -1)) != int(epoch)53 ):54 return None55 state = deepcopy(self._pending_resume_state)56 self._pending_resume_state = None57 return state58 59 @staticmethod60 def _validate_resume_topology(61 resume_state: dict,62 *,63 context: SourceContext,64 loader_num_workers: int,65 ) -> None:66 resume_topology = resume_state.get(_RESUME_TOPOLOGY_KEY)67 if not isinstance(resume_topology, dict):68 raise RuntimeError(69 "Resume state is missing required worker topology metadata."70 )71 expected_world_size = int(resume_topology["world_size"])72 expected_num_workers = int(resume_topology["loader_num_workers"])73 expected_global_worker_count = int(resume_topology["global_worker_count"])74 current_num_workers = int(loader_num_workers)75 current_global_worker_count = int(context.global_worker_count)76 if (77 expected_world_size != int(context.world_size)78 or expected_num_workers != current_num_workers79 or expected_global_worker_count != current_global_worker_count80 ):81 raise RuntimeError(82 "Resume requires the same data worker topology as the saved state. "83 f"saved(world_size={expected_world_size}, "84 f"num_workers_per_rank={expected_num_workers}, "85 f"global_worker_count={expected_global_worker_count}), "86 f"current(world_size={context.world_size}, "87 f"num_workers_per_rank={current_num_workers}, "88 f"global_worker_count={current_global_worker_count})."89 )90 91 def __iter__(self) -> Iterable[dict]:92 worker_info = get_worker_info()93 if worker_info is None:94 worker_id = 095 loader_num_workers = 096 effective_num_workers = 197 else:98 worker_id = worker_info.id99 loader_num_workers = worker_info.num_workers100 effective_num_workers = worker_info.num_workers101 102 epoch = self._current_epoch()103 context = SourceContext(104 epoch=epoch,105 rank=self.rank,106 world_size=self.world_size,107 worker_id=worker_id,108 num_workers=effective_num_workers,109 seed=self.seed,110 )111 resume_state = self._take_resume_state(epoch)112 if resume_state is not None:113 self._validate_resume_topology(114 resume_state,115 context=context,116 loader_num_workers=loader_num_workers,117 )118 worker_state = (119 None120 if resume_state is None121 else (resume_state.get("workers") or {}).get(str(context.global_worker_id))122 )123 sample_iter = self.source.iter_samples(124 context,125 state=None if worker_state is None else worker_state.get("adapter_state"),126 )127 for sample in sample_iter:128 sample["data_worker_id"] = context.worker_id129 sample["data_global_worker_id"] = context.global_worker_id130 yield sample131 132 133class _DataStateTracker:134 def __init__(self, *, num_tokens_per_epoch: int | None):135 self.num_tokens_per_epoch = (136 None if num_tokens_per_epoch is None else int(num_tokens_per_epoch)137 )138 self._pending_state: dict | None = None139 self._reset_for_epoch(epoch=0)140 141 def _reset_for_epoch(self, *, epoch: int) -> None:142 self.epoch = int(epoch)143 self.samples_emitted = 0144 self.num_text_tokens = 0145 self.num_audio_tokens = 0146 self.num_total_tokens = 0147 self.workers: dict[str, dict] = {}148 self._next_sample_order_by_worker: dict[str, int] = {}149 150 def load_state_dict(self, state: dict | None) -> None:151 self._pending_state = deepcopy(state) if state else None152 153 def set_epoch(self, epoch: int) -> None:154 if self._pending_state is not None and int(155 self._pending_state.get("epoch", -1)156 ) == int(epoch):157 state = deepcopy(self._pending_state)158 self._pending_state = None159 self.epoch = int(state.get("epoch", epoch))160 self.samples_emitted = int(state.get("samples_emitted", 0))161 self.num_text_tokens = int(state.get("num_text_tokens", 0))162 self.num_audio_tokens = int(state.get("num_audio_tokens", 0))163 self.num_total_tokens = int(state.get("num_total_tokens", 0))164 self.workers = deepcopy(state.get("workers") or {})165 self._next_sample_order_by_worker = {166 worker_key: int((worker_state or {}).get("sample_order", -1)) + 1167 for worker_key, worker_state in self.workers.items()168 }169 return170 self._reset_for_epoch(epoch=int(epoch))171 172 def should_stop(self) -> bool:173 return (174 self.num_tokens_per_epoch is not None175 and self.num_total_tokens >= self.num_tokens_per_epoch176 )177 178 def stage_sample(self, sample: dict) -> dict:179 item = dict(sample)180 worker_key = str(item.pop("data_global_worker_id"))181 item.pop("data_worker_id", None)182 adapter_state = item.pop("_adapter_state", None)183 sample_order = int(self._next_sample_order_by_worker.get(worker_key, 0))184 self._next_sample_order_by_worker[worker_key] = sample_order + 1185 item[_TRACKING_KEY] = {186 "worker_key": worker_key,187 "adapter_state": deepcopy(adapter_state),188 "sample_order": sample_order,189 "num_text_tokens": int(item["num_text_tokens"]),190 "num_audio_tokens": int(item["num_audio_tokens"]),191 "num_total_tokens": int(192 item.get("num_total_tokens", item["input_ids_length"])193 ),194 }195 return item196 197 def _pop_tracking(self, sample: dict) -> tuple[dict, dict]:198 item = dict(sample)199 tracking = item.pop(_TRACKING_KEY, None)200 if not isinstance(tracking, dict):201 raise RuntimeError("Tracked sample is missing internal resume metadata.")202 return item, tracking203 204 def _advance_worker(self, tracking: dict) -> None:205 adapter_state = tracking.get("adapter_state")206 if adapter_state is None:207 return208 worker_key = str(tracking["worker_key"])209 sample_order = int(tracking.get("sample_order", -1))210 current_state = self.workers.get(worker_key)211 current_order = int((current_state or {}).get("sample_order", -1))212 if current_order >= sample_order:213 return214 self.workers[worker_key] = {215 "adapter_state": deepcopy(adapter_state),216 "sample_order": sample_order,217 }218 219 def mark_samples_dropped(self, samples: list[dict]) -> None:220 for sample in samples:221 _, tracking = self._pop_tracking(sample)222 self._advance_worker(tracking)223 224 def commit_batch(self, samples: list[dict]) -> list[dict]:225 committed: list[dict] = []226 for sample in samples:227 item, tracking = self._pop_tracking(sample)228 self._advance_worker(tracking)229 self.samples_emitted += 1230 self.num_text_tokens += int(tracking["num_text_tokens"])231 self.num_audio_tokens += int(tracking["num_audio_tokens"])232 self.num_total_tokens += int(tracking["num_total_tokens"])233 committed.append(item)234 return committed235 236 def state_dict(self) -> dict:237 return {238 "epoch": int(self.epoch),239 "samples_emitted": int(self.samples_emitted),240 "num_text_tokens": int(self.num_text_tokens),241 "num_audio_tokens": int(self.num_audio_tokens),242 "num_total_tokens": int(self.num_total_tokens),243 "workers": deepcopy(self.workers),244 "num_tokens_per_epoch": self.num_tokens_per_epoch,245 }246 247 248class BatchedDataStream:249 def __init__(250 self,251 *,252 sample_dataset: StreamingSampleDataset,253 data_cfg,254 tokenizer,255 num_tokens_per_epoch: int | None,256 profiler=None,257 ):258 from dots_tts.data.collator import PadCollator259 260 self.sample_dataset = sample_dataset261 self.profiler = ensure_data_profiler(profiler)262 llm_token_rate = (263 float(data_cfg.train_audio_sample_rate)264 / float(data_cfg.audio_samples_per_llm_token)265 )266 self.batcher = OnlineBatcher(267 max_audio_tokens_in_batch=max(268 1,269 math.ceil(float(data_cfg.max_audio_seconds_in_batch) * llm_token_rate),270 ),271 max_text_tokens_in_batch=data_cfg.max_text_tokens_in_batch,272 max_batch_size=data_cfg.max_samples_per_batch,273 sample_pool_size=data_cfg.bucketing_pool_size,274 profiler=self.profiler,275 )276 self.sample_loader = None277 self.collator = PadCollator(tokenizer)278 self.data_state = _DataStateTracker(279 num_tokens_per_epoch=num_tokens_per_epoch280 )281 self._decision_iterator = None282 self._sample_iterator = None283 self._pending_batch = None284 self._pending_samples = None285 286 def attach_loader(self, loader: DataLoader) -> None:287 self.sample_loader = loader288 289 def close(self) -> None:290 self._reset_iteration_state()291 self.sample_loader = None292 293 def load_state_dict(self, state: dict | None) -> None:294 self.data_state.load_state_dict(state)295 self.sample_dataset.load_state_dict(state)296 self._reset_iteration_state()297 298 def state_dict(self) -> dict:299 if self.sample_loader is None:300 raise RuntimeError("BatchedDataStream has no attached sample loader.")301 if self._pending_batch is not None or self._pending_samples is not None:302 raise RuntimeError(303 "Cannot serialize BatchedDataStream while a batch is pending commit."304 )305 loader_num_workers = int(getattr(self.sample_loader, "num_workers", 0))306 effective_num_workers = max(1, loader_num_workers)307 state = self.data_state.state_dict()308 state[_RESUME_TOPOLOGY_KEY] = {309 "world_size": int(self.sample_dataset.world_size),310 "loader_num_workers": loader_num_workers,311 "global_worker_count": int(self.sample_dataset.world_size)312 * effective_num_workers,313 }314 return state315 316 def set_epoch(self, epoch: int) -> None:317 self.sample_dataset.set_epoch(epoch)318 self.data_state.set_epoch(epoch)319 self._reset_iteration_state()320 321 def _reset_iteration_state(self) -> None:322 close_iterator = getattr(self._decision_iterator, "close", None)323 if callable(close_iterator):324 close_iterator()325 self._decision_iterator = None326 self._sample_iterator = None327 self._pending_batch = None328 self._pending_samples = None329 330 def _iter_staged_samples(self):331 if self.sample_loader is None:332 raise RuntimeError("BatchedDataStream has no attached sample loader.")333 self._sample_iterator = iter(self.sample_loader)334 profiler = self.profiler335 try:336 while True:337 if self.data_state.should_stop():338 return339 try:340 with profiler.measure("main.loader_wait_next_sample"):341 sample = next(self._sample_iterator)342 except StopIteration:343 return344 if sample is None:345 continue346 with profiler.measure("main.stage_sample"):347 staged = self.data_state.stage_sample(sample)348 yield staged349 finally:350 self._sample_iterator = None351 352 def _decision_stream(self):353 if self._decision_iterator is None:354 self._decision_iterator = iter(355 self.batcher.build_decisions(self._iter_staged_samples())356 )357 return self._decision_iterator358 359 def peek_batch(self) -> tuple[dict | None, bool]:360 if self._pending_batch is not None:361 return self._pending_batch, True362 363 for decision in self._decision_stream():364 if decision.dropped_samples:365 self.data_state.mark_samples_dropped(decision.dropped_samples)366 if not decision.batch_samples:367 continue368 self._pending_samples = decision.batch_samples369 with self.profiler.measure(370 "main.collate_batch",371 count=len(decision.batch_samples),372 ):373 self._pending_batch = self.collator(decision.batch_samples)374 return self._pending_batch, True375 return None, False376 377 def commit_batch(self) -> dict:378 if self._pending_batch is None or self._pending_samples is None:379 raise RuntimeError("BatchedDataStream has no pending batch to commit.")380 pending_batch = self._pending_batch381 self.data_state.commit_batch(self._pending_samples)382 self._pending_batch = None383 self._pending_samples = None384 return pending_batch385 386 def discard_batch(self) -> None:387 if self._pending_batch is None or self._pending_samples is None:388 raise RuntimeError("BatchedDataStream has no pending batch to discard.")389 self._pending_batch = None390 self._pending_samples = None391 392 def __iter__(self):393 while True:394 batch, has_batch = self.peek_batch()395 if not has_batch:396 return397 self.commit_batch()398 yield batch399 if self.data_state.should_stop():400 return401 