195 lines
5.5 KiB
Python
195 lines
5.5 KiB
Python
from __future__ import annotations
|
|
|
|
from torch.utils.data import DataLoader
|
|
|
|
from dots_tts.config.data import DataConfig
|
|
from dots_tts.data.pipelines.base import BaseSamplePipeline
|
|
from dots_tts.data.pipelines.tts_pipeline import BasicTtsPipeline, InterleaveTtsPipeline
|
|
from dots_tts.data.source_adapters.jsonl_manifest_adapter import (
|
|
JsonlManifestSourceAdapter,
|
|
)
|
|
from dots_tts.data.source_adapters.multi_source_adapter import (
|
|
SequentialMultiSourceAdapter,
|
|
SourceSpec,
|
|
WeightedMultiSourceAdapter,
|
|
)
|
|
from dots_tts.data.streaming import (
|
|
BatchedDataStream,
|
|
StreamingSampleDataset,
|
|
identity_collate,
|
|
)
|
|
|
|
_SOURCE_ADAPTER_CLASSES = {
|
|
"JsonlManifestSourceAdapter": JsonlManifestSourceAdapter,
|
|
}
|
|
|
|
|
|
def _build_source_pipeline(
|
|
tokenizer, data_cfg, pipeline_name: str, *, profiler=None
|
|
) -> BaseSamplePipeline:
|
|
if pipeline_name == "basic":
|
|
return BasicTtsPipeline(tokenizer, data_cfg, profiler=profiler)
|
|
if pipeline_name == "interleave":
|
|
return InterleaveTtsPipeline(tokenizer, data_cfg, profiler=profiler)
|
|
raise ValueError(f"Unsupported data pipeline: {pipeline_name!r}")
|
|
|
|
|
|
def _build_source_specs(data_cfg, tokenizer, *, profiler=None) -> list[SourceSpec]:
|
|
specs = []
|
|
for source_cfg in data_cfg.sources:
|
|
adapter_cls = _SOURCE_ADAPTER_CLASSES[source_cfg.adapter.class_name]
|
|
adapter = adapter_cls(**source_cfg.adapter.params)
|
|
specs.append(
|
|
SourceSpec(
|
|
name=source_cfg.name,
|
|
weight=float(source_cfg.weight),
|
|
adapter=adapter,
|
|
pipeline=_build_source_pipeline(
|
|
tokenizer, data_cfg, source_cfg.pipeline, profiler=profiler
|
|
),
|
|
)
|
|
)
|
|
return specs
|
|
|
|
|
|
def _resolve_rank_info(accelerator=None) -> tuple[int, int]:
|
|
rank = (
|
|
int(getattr(accelerator, "process_index", 0)) if accelerator is not None else 0
|
|
)
|
|
world_size = (
|
|
int(getattr(accelerator, "num_processes", 1)) if accelerator is not None else 1
|
|
)
|
|
return rank, world_size
|
|
|
|
|
|
def _local_num_tokens_per_epoch(
|
|
global_num_tokens_per_epoch: int, *, rank: int, world_size: int
|
|
) -> int:
|
|
if world_size <= 0:
|
|
raise ValueError(f"world_size must be positive, but got {world_size}.")
|
|
if rank < 0 or rank >= world_size:
|
|
raise ValueError(
|
|
f"rank must be in [0, {world_size}), but got rank={rank}."
|
|
)
|
|
|
|
base, remainder = divmod(int(global_num_tokens_per_epoch), int(world_size))
|
|
return base + int(rank < remainder)
|
|
|
|
|
|
def _build_dataset(
|
|
data_cfg: DataConfig,
|
|
*,
|
|
tokenizer,
|
|
seed: int,
|
|
accelerator=None,
|
|
sequential: bool,
|
|
profiler=None,
|
|
):
|
|
rank, world_size = _resolve_rank_info(accelerator)
|
|
source_cls = SequentialMultiSourceAdapter if sequential else WeightedMultiSourceAdapter
|
|
source = source_cls(
|
|
sources=_build_source_specs(data_cfg, tokenizer, profiler=profiler)
|
|
)
|
|
return StreamingSampleDataset(
|
|
source=source,
|
|
rank=rank,
|
|
world_size=world_size,
|
|
seed=int(seed),
|
|
)
|
|
|
|
|
|
def build_training_dataset(
|
|
data_cfg: DataConfig,
|
|
tokenizer,
|
|
*,
|
|
seed: int,
|
|
accelerator=None,
|
|
profiler=None,
|
|
):
|
|
if data_cfg.num_tokens_per_epoch is None:
|
|
raise ValueError("Training data requires num_tokens_per_epoch.")
|
|
return _build_dataset(
|
|
data_cfg,
|
|
tokenizer=tokenizer,
|
|
seed=seed,
|
|
accelerator=accelerator,
|
|
sequential=False,
|
|
profiler=profiler,
|
|
)
|
|
|
|
|
|
def build_validation_dataset(
|
|
data_cfg: DataConfig,
|
|
tokenizer,
|
|
*,
|
|
seed: int,
|
|
accelerator=None,
|
|
profiler=None,
|
|
):
|
|
return _build_dataset(
|
|
data_cfg,
|
|
tokenizer=tokenizer,
|
|
seed=seed,
|
|
accelerator=accelerator,
|
|
sequential=True,
|
|
profiler=profiler,
|
|
)
|
|
|
|
|
|
def _build_sample_loader(dataset, data_cfg: DataConfig) -> DataLoader:
|
|
loader_kwargs = {
|
|
"dataset": dataset,
|
|
"batch_size": None,
|
|
"collate_fn": identity_collate,
|
|
"num_workers": data_cfg.num_workers,
|
|
"pin_memory": data_cfg.pin_memory,
|
|
"persistent_workers": data_cfg.num_workers > 0,
|
|
}
|
|
if data_cfg.num_workers > 0:
|
|
loader_kwargs["prefetch_factor"] = int(data_cfg.prefetch_factor)
|
|
sample_loader = DataLoader(**loader_kwargs)
|
|
return sample_loader
|
|
|
|
|
|
def build_training_dataloader(
|
|
dataset, data_cfg: DataConfig, tokenizer, *, profiler=None
|
|
):
|
|
local_num_tokens_per_epoch = _local_num_tokens_per_epoch(
|
|
int(data_cfg.num_tokens_per_epoch),
|
|
rank=int(dataset.rank),
|
|
world_size=int(dataset.world_size),
|
|
)
|
|
sample_loader = _build_sample_loader(dataset, data_cfg)
|
|
batched_stream = BatchedDataStream(
|
|
sample_dataset=dataset,
|
|
data_cfg=data_cfg,
|
|
tokenizer=tokenizer,
|
|
num_tokens_per_epoch=local_num_tokens_per_epoch,
|
|
profiler=profiler,
|
|
)
|
|
batched_stream.attach_loader(sample_loader)
|
|
return batched_stream
|
|
|
|
|
|
def build_validation_dataloader(
|
|
dataset, data_cfg: DataConfig, tokenizer, *, profiler=None
|
|
):
|
|
sample_loader = _build_sample_loader(dataset, data_cfg)
|
|
batched_stream = BatchedDataStream(
|
|
sample_dataset=dataset,
|
|
data_cfg=data_cfg,
|
|
tokenizer=tokenizer,
|
|
num_tokens_per_epoch=None,
|
|
profiler=profiler,
|
|
)
|
|
batched_stream.attach_loader(sample_loader)
|
|
return batched_stream
|
|
|
|
|
|
__all__ = [
|
|
"build_training_dataloader",
|
|
"build_training_dataset",
|
|
"build_validation_dataloader",
|
|
"build_validation_dataset",
|
|
]
|