MCPcopy Create free account
hub / github.com/FireRedTeam/FireRedTTS2 / __init__

Method __init__

fireredtts2/codec/rvq.py:17–54  ·  view source on GitHub ↗
(
        self,
        input_dim: int,
        codebook_size: int,
        codebook_dim: int,
    )

Source from the content-addressed store, hash-verified

15
16class VectorQuantize(nn.Module):
17 def __init__(
18 self,
19 input_dim: int,
20 codebook_size: int,
21 codebook_dim: int,
22 ):
23 super().__init__()
24 self.input_dim = input_dim
25 self.codebook_size = codebook_size
26 self.codebook_dim = codebook_dim
27
28 self.in_project = (
29 WNConv1d(
30 self.input_dim, self.codebook_dim, kernel_size=1
31 ) # (B, D, T) -> (B, D', T)
32 if self.input_dim != self.codebook_dim
33 else nn.Identity()
34 )
35 self.out_project = (
36 WNConv1d(
37 self.codebook_dim, self.input_dim, kernel_size=1
38 ) # (B, D', T) -> (B, D, T)
39 if self.input_dim != self.codebook_dim
40 else nn.Identity()
41 )
42
43 # Initialize codebook and EMA buffers
44 self.register_buffer(
45 "codebook", torch.zeros(codebook_size, codebook_dim).float()
46 ) # (codebook_size, D'), ensure fp32
47 # Place holder, not used in inference
48 self.register_buffer("inited", torch.tensor([True], dtype=torch.bool)) # (1)
49 self.register_buffer(
50 "cluster_size", torch.zeros(codebook_size).float()
51 ) # (codebook_size), ensure fp32
52 self.register_buffer(
53 "embed_avg", self.codebook.clone().float()
54 ) # (codebook_size, D'), ensure fp32
55
56 def decode_code(self, embed_id): # embed_id: (B, T)
57 embed = (

Callers

nothing calls this directly

Calls 2

WNConv1dFunction · 0.85
__init__Method · 0.45

Tested by

no test coverage detected