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

Class VectorQuantize

fireredtts2/codec/rvq.py:16–89  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

14
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 = (
58 F.embedding(embed_id, self.codebook).transpose(1, 2).float()
59 ) # (B, D', T), ensure fp32
60 return embed
61
62 def encode_code(self, z: torch.Tensor): # z: (B, D, T)
63 # logging.info(f"{self.cluster_size = }, {self.codebook = }, {self.embed_avg = }, {self.inited = }")
64 z = z.float() # Ensure fp32
65 z_e = self.in_project(z).float() # (B, D', T), ensure fp32
66
67 # Rearrange for quantization
68 encodings = rearrange(z_e, "b d t -> (b t) d").float() # (B*T, D'), ensure fp32
69
70 # Quantization
71 dist = (
72 encodings.pow(2).sum(1, keepdim=True) # (B*T, 1)
73 - 2 * encodings @ self.codebook.float().t() # (B*T, codebook_size)

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected