File size: 2,732 Bytes
98af51e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
"""
WavCoch configuration for Hugging Face Transformers.
"""

from transformers import PretrainedConfig


class WavCochConfig(PretrainedConfig):
    """Configuration class for WavCoch checkpoints with optional vocoder."""

    model_type = "wavcoch"

    def __init__(
        self,
        window_size: int = 1001,
        window_padding: int = 1000,
        hop_length: int = 80,
        out_channels: int = 211,
        causal_convs: bool = True,
        causal_pad_mode: str = "repeat",
        encoder_layers: int = 8,
        encoder_dim: int = 512,
        encoder_kernel_size: int = 3,
        decoder_layers: int = 8,
        decoder_dim: int = 512,
        decoder_kernel_size: int = 9,
        quantizer: str = "FSQ",
        channels=None,
        vocab_size: int = None,
        sample_rate: int = 16000,
        has_vocoder: bool = False,
        vocoder_upsample_rates=None,
        vocoder_upsample_kernel_sizes=None,
        vocoder_upsample_initial_channel: int = 512,
        vocoder_resblock: str = "1",
        vocoder_resblock_kernel_sizes=None,
        vocoder_resblock_dilation_sizes=None,
        **kwargs,
    ):
        channels = list(channels or [8, 8, 8, 4, 4])
        if vocab_size is None:
            vocab_size = 1
            for level in channels:
                vocab_size *= int(level)

        self.window_size = int(window_size)
        self.window_padding = int(window_padding)
        self.hop_length = int(hop_length)
        self.out_channels = int(out_channels)
        self.causal_convs = bool(causal_convs)
        self.causal_pad_mode = str(causal_pad_mode)
        self.encoder_layers = int(encoder_layers)
        self.encoder_dim = int(encoder_dim)
        self.encoder_kernel_size = int(encoder_kernel_size)
        self.decoder_layers = int(decoder_layers)
        self.decoder_dim = int(decoder_dim)
        self.decoder_kernel_size = int(decoder_kernel_size)
        self.quantizer = str(quantizer)
        self.channels = channels
        self.vocab_size = int(vocab_size)
        self.sample_rate = int(sample_rate)

        self.has_vocoder = bool(has_vocoder)
        self.vocoder_upsample_rates = list(vocoder_upsample_rates or [5, 4, 2, 2])
        self.vocoder_upsample_kernel_sizes = list(vocoder_upsample_kernel_sizes or [10, 8, 4, 4])
        self.vocoder_upsample_initial_channel = int(vocoder_upsample_initial_channel)
        self.vocoder_resblock = str(vocoder_resblock)
        self.vocoder_resblock_kernel_sizes = list(vocoder_resblock_kernel_sizes or [11, 7, 3])
        self.vocoder_resblock_dilation_sizes = [
            list(d) for d in (vocoder_resblock_dilation_sizes or [[1, 3, 5], [1, 3, 5], [1, 3, 5]])
        ]

        super().__init__(**kwargs)