CoolFace
Apppublic

RustyMark/dots.tts

sourceHugging Faceapache-2.0updated 3mo agoView on Hugging Face
0likes
streaming.py401 linesDownload Raw Back to data
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