(
self,
dimension: int = 256,
n_q: int = 8,
bins: int = 1024,
decay: float = 0.99,
kmeans_init: bool = True,
kmeans_iters: int = 50,
threshold_ema_dead_code: int = 2,
)
| 40 | """ |
| 41 | |
| 42 | def __init__( |
| 43 | self, |
| 44 | dimension: int = 256, |
| 45 | n_q: int = 8, |
| 46 | bins: int = 1024, |
| 47 | decay: float = 0.99, |
| 48 | kmeans_init: bool = True, |
| 49 | kmeans_iters: int = 50, |
| 50 | threshold_ema_dead_code: int = 2, |
| 51 | ): |
| 52 | super().__init__() |
| 53 | self.n_q = n_q |
| 54 | self.dimension = dimension |
| 55 | self.bins = bins |
| 56 | self.decay = decay |
| 57 | self.kmeans_init = kmeans_init |
| 58 | self.kmeans_iters = kmeans_iters |
| 59 | self.threshold_ema_dead_code = threshold_ema_dead_code |
| 60 | self.vq = ResidualVectorQuantization( |
| 61 | dim=self.dimension, |
| 62 | codebook_size=self.bins, |
| 63 | num_quantizers=self.n_q, |
| 64 | decay=self.decay, |
| 65 | kmeans_init=self.kmeans_init, |
| 66 | kmeans_iters=self.kmeans_iters, |
| 67 | threshold_ema_dead_code=self.threshold_ema_dead_code, |
| 68 | ) |
| 69 | |
| 70 | def forward( |
| 71 | self, |
nothing calls this directly
no test coverage detected