RustyMark/dots.tts
0
1from __future__ import annotations2 3from typing import Any4 5import torch6from torch.nn.utils.rnn import pad_sequence7 8 9class PadCollator:10 def __init__(self, tokenizer):11 self.tokenizer = tokenizer12 self.pad_token_id = tokenizer.pad_token_id13 if self.pad_token_id is None:14 self.pad_token_id = tokenizer.eos_token_id or 015 16 def __call__(self, samples: list[dict[str, Any]]) -> dict[str, Any]:17 if not samples:18 raise ValueError("PadCollator received an empty sample list.")19 20 order = sorted(21 range(len(samples)),22 key=lambda idx: samples[idx]["sample_length"],23 reverse=True,24 )25 ordered = [samples[idx] for idx in order]26 27 input_ids = [28 torch.tensor(sample["input_ids"], dtype=torch.long) for sample in ordered29 ]30 labels = [31 torch.tensor(sample["labels"], dtype=torch.long) for sample in ordered32 ]33 loss_masks = [34 torch.tensor(sample["loss_mask"], dtype=torch.float32) for sample in ordered35 ]36 waveforms = [sample["sample"].squeeze(0) for sample in ordered]37 fbank = [sample["fbank"] for sample in ordered]38 39 return {40 "fids": [sample["fid"] for sample in ordered],41 "source_names": [sample.get("source_name") for sample in ordered],42 "input_ids": pad_sequence(43 input_ids,44 batch_first=True,45 padding_value=self.pad_token_id,46 ),47 "input_ids_lengths": torch.tensor(48 [len(sample["input_ids"]) for sample in ordered],49 dtype=torch.long,50 ),51 "labels": pad_sequence(52 labels,53 batch_first=True,54 padding_value=self.pad_token_id,55 ),56 "loss_mask": pad_sequence(57 loss_masks,58 batch_first=True,59 padding_value=0.0,60 ),61 "sample": pad_sequence(62 waveforms,63 batch_first=True,64 padding_value=0.0,65 ).unsqueeze(1),66 "sample_lengths": torch.tensor(67 [sample["sample_length"] for sample in ordered],68 dtype=torch.long,69 ),70 "num_text_tokens": torch.tensor(71 [sample["num_text_tokens"] for sample in ordered],72 dtype=torch.long,73 ),74 "num_audio_tokens": torch.tensor(75 [sample["num_audio_tokens"] for sample in ordered],76 dtype=torch.long,77 ),78 "fbank": pad_sequence(79 fbank,80 batch_first=True,81 padding_value=0.0,82 ),83 "fbank_lengths": torch.tensor(84 [sample["fbank_length"] for sample in ordered],85 dtype=torch.long,86 ),87 }88 