espnet2.tasks.rst.RestorationCollateFn
Less than 1 minute
espnet2.tasks.rst.RestorationCollateFn
class espnet2.tasks.rst.RestorationCollateFn(max_samples: int, input_sr: int, noise_dir: str, rir_dir: str, degrade_prob: float, online_degradation: bool, train: bool = True)
Bases: object
Collate function for restoration training: online degradation + padding.
Not specific to one method: any SSL encoder / vocoder pair trains with it. The degradation pipeline itself follows the Sidon paper (see degrade_waveform).
SSL feature extraction is done on GPU in the model forward pass (W2VBert2Encoder._wav_to_ssl_inputs), not here.
