class
ResidualVectorQuantizer
PythonResidual Vector Quantizer (RVQ) with straight-through estimator.
ResidualVectorQuantizer( num_levels: int, num_embeddings: int | Sequence[int], embedding_dim: int, beta: float = 0.25, **kwargs={},)Residual Vector Quantizer (RVQ) with straight-through estimator.
Input: […, D] (last dim = embedding_dim) Output: […, D] (sum of per-level dequantized vectors; gradients pass through x)
Metrics (logged via metrics property):
- rvq_l{l}_perplexity, rvq_l{l}_usage, rvq_l{l}_bits_per_index
- rvq_perplexity_mean, rvq_usage_mean, rvq_bits_per_index_sum (entropy lower bound)
Parameters
| Name | Type | Default | Description |
|---|---|---|---|
num_levels | int | Required | int, number of residual VQ stages (M >= 1) |
num_embeddings | int | Sequence[int] | Required | int OR sequence[int], codebook size K for each level |
embedding_dim | int | Required | int, latent dimensionality D |
beta | float | 0.25 | float, commitment coefficient per level |
constant
M
PythonM = int(num_levels)constant
D
PythonD = int(embedding_dim)attribute
Ks
PythonKsattribute
beta
Pythonbeta = float(beta)attribute
metrics
Pythonmetricsmethod
build
Pythonbuild(input_shape)method
call
Pythoncall( x: keras.KerasTensor, return_indices: bool = False,) -> keras.KerasTensor | tuple[keras.KerasTensor, list[keras.KerasTensor]]Parameters
| Name | Type | Default | Description |
|---|---|---|---|
x | keras.KerasTensor | Required | [..., D] latent to be quantized. |
return_indices | bool | False | if True, also returns list of flat indices (one tensor per level). |
Returns
| Type | Description |
|---|---|
keras.KerasTensor | tuple[keras.KerasTensor, list[keras.KerasTensor]] | y or (y, indices_list): dequantized vector and optional per-level indices. |
method
encode
PythonReturn list of per-level flat index tensors [N] (no gradients).
encode(x)Return list of per-level flat index tensors [N] (no gradients).
method
decode
PythonSum per-level code vectors from indiceslist and reshape to originalshape.
decode(indices_list, original_shape)Sum per-level code vectors from indices_list and reshape to original_shape.
method
get_config
Pythonget_config()