Encode a given input tensor with the specified sample rate at the given bandwidth. The RVQ encode method sets the appropriate number of quantizer to use and returns indices for each quantizer. Args: x (torch.Tensor): Input tensor. n_q (int): Number of
(
self, x: torch.Tensor, n_q: tp.Optional[int] = None, st: tp.Optional[int] = None
)
| 94 | return quantized, codes, torch.mean(commit_loss), quantized_list |
| 95 | |
| 96 | def encode( |
| 97 | self, x: torch.Tensor, n_q: tp.Optional[int] = None, st: tp.Optional[int] = None |
| 98 | ) -> torch.Tensor: |
| 99 | """Encode a given input tensor with the specified sample rate at the given bandwidth. |
| 100 | The RVQ encode method sets the appropriate number of quantizer to use |
| 101 | and returns indices for each quantizer. |
| 102 | Args: |
| 103 | x (torch.Tensor): Input tensor. |
| 104 | n_q (int): Number of quantizer used to quantize. Default: All quantizers. |
| 105 | st (int): Start to encode input from which layers. Default: 0. |
| 106 | """ |
| 107 | n_q = n_q if n_q else self.n_q |
| 108 | st = st or 0 |
| 109 | codes = self.vq.encode(x, n_q=n_q, st=st) |
| 110 | return codes |
| 111 | |
| 112 | def decode(self, codes: torch.Tensor, st: int = 0) -> torch.Tensor: |
| 113 | """Decode the given codes to the quantized representation. |
nothing calls this directly
no outgoing calls
no test coverage detected