MCPcopy Create free account
hub / github.com/1038lab/ComfyUI-FireRedTTS / encode_code

Method encode_code

fireredtts2/codec/rvq.py:62–89  ·  view source on GitHub ↗
(self, z: torch.Tensor)

Source from the content-addressed store, hash-verified

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)
74 + self.codebook.float().pow(2).sum(1, keepdim=True).t()
75 ) # (1, codebook_size)
76
77 # dist: (B*T, codebook_size)
78 indices = (-dist).max(1)[1] # (B*T)
79 indices = rearrange(indices, "(b t) -> b t", b=z.size(0)) # (B, T)
80
81 # Get quantized vectors
82 z_q = self.decode_code(indices).float() # (B, D', T), ensure fp32
83
84 # Straight-through estimator
85 z_q = z_e + (z_q - z_e).detach() # (B, D', T)
86 z_q = self.out_project(z_q).float() # (B, D, T), ensure fp32
87
88 # z_q: (B, D, T), commit_loss: (B), indices: (B, T), z: (B, D', T)
89 return z_q, indices
90
91
92class ResidualVQ(nn.Module):

Callers 1

encode_codesMethod · 0.80

Calls 1

decode_codeMethod · 0.95

Tested by

no test coverage detected