Residual VQ with EMA codebook updates.
EmaResidualVectorQuantizer( num_levels: int, num_embeddings: int | Sequence[int], embedding_dim: int, beta: float = 0.25, ema_decay: float = 0.99, epsilon: float = 1e-05, **kwargs={},)Residual VQ with EMA codebook updates.
Instead of learning codebook embeddings via gradient descent (which requires a codebook loss term), this layer maintains exponential moving averages of cluster assignment counts and embedding sums. Codebook vectors are derived from these running statistics with Laplace smoothing for numerical stability.
Only the commitment loss is back-propagated through the encoder; the straight-through estimator copies gradients from the decoder to the encoder as in the standard VQ-VAE.
Input: [..., D] (last dim = embedding_dim)
Output: [..., D] (sum of per-level dequantized vectors)
Metrics (logged via metrics property):
rvq_l{l}_perplexity,rvq_l{l}_usage,rvq_l{l}_bits_per_indexrvq_perplexity_mean,rvq_usage_mean,rvq_bits_per_index_sum(entropy lower bound)
Example:
rvq = EmaResidualVectorQuantizer( num_levels=4, num_embeddings=64, embedding_dim=16, ema_decay=0.99,)y = rvq(z, training=True) # forward + EMA updatey, indices = rvq(z, return_indices=True) # also return codesParameters
| Name | Type | Default | Description |
|---|---|---|---|
num_levels | int | Required | Number of residual VQ stages (``M >= 1``). |
num_embeddings | int | Sequence[int] | Required | Codebook size ``K`` per level (int or per-level list). |
embedding_dim | int | Required | Latent dimensionality ``D``. |
beta | float | 0.25 | Commitment loss coefficient. |
ema_decay | float | 0.99 | EMA decay rate for codebook updates (``0.99``–``0.999`` typical). |
epsilon | float | 1e-05 | Small constant for Laplace smoothing of cluster counts. |
M
PythonM = int(num_levels)D
PythonD = int(embedding_dim)Ks
PythonKsbeta
Pythonbeta = float(beta)ema_decay
Pythonema_decay = float(ema_decay)epsilon
Pythonepsilon = float(epsilon)metrics
PythonExpose per-level + aggregate metrics so Model.fit logs them.
metricsExpose per-level + aggregate metrics so Model.fit logs them.
build
Pythonbuild(input_shape)call
PythonQuantize x through all residual levels.
call( x: keras.KerasTensor, training: bool = False, return_indices: bool = False,) -> keras.KerasTensor | tuple[keras.KerasTensor, list[keras.KerasTensor]]Quantize x through all residual levels.
Parameters
| Name | Type | Default | Description |
|---|---|---|---|
x | keras.KerasTensor | Required | ``[..., D]`` latent to be quantized. |
training | bool | False | If ``True``, run EMA codebook updates. |
return_indices | bool | False | If ``True``, also return per-level flat indices. |
Returns
| Type | Description |
|---|---|
keras.KerasTensor | tuple[keras.KerasTensor, list[keras.KerasTensor]] | ``y`` or ``(y, indices_list)``: dequantized vector and optional |
keras.KerasTensor | tuple[keras.KerasTensor, list[keras.KerasTensor]] | per-level index tensors. |
encode
PythonReturn list of per-level flat index tensors [N] (no gradients).
encode(x: keras.KerasTensor) -> list[keras.KerasTensor]Return list of per-level flat index tensors [N] (no gradients).
decode
PythonSum per-level code vectors from indiceslist and reshape.
decode(indices_list: list[keras.KerasTensor], original_shape: tuple[int, ...]) -> keras.KerasTensorSum per-level code vectors from indices_list and reshape.
get_config
Pythonget_config()