CoolFace
Apppublic

RustyMark/dots.tts

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