Skip to content
heliaEDGE
Reference
HELIA

ema_residual_vector_quantizer

Residual Vector Quantizer with Exponential Moving Average codebook updates.

Replaces the gradient-based codebook loss with EMA updates to codebook embeddings, following van den Oord et al. 2017 (VQ-VAE). Only the commitment loss is back-propagated; codebook vectors are updated via running averages of assigned encoder outputs.

Machine-readable model

class

Residual VQ with EMA codebook updates.

helia_edge/layers/ema_residual_vector_quantizer.py:18

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_index
  • rvq_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 update
y, indices = rvq(z, return_indices=True) # also return codes
Parameters of EmaResidualVectorQuantizer
NameTypeDefaultDescription
num_levelsintRequiredNumber of residual VQ stages (``M >= 1``).
num_embeddingsint | Sequence[int]RequiredCodebook size ``K`` per level (int or per-level list).
embedding_dimintRequiredLatent dimensionality ``D``.
betafloat0.25Commitment loss coefficient.
ema_decayfloat0.99EMA decay rate for codebook updates (``0.99``–``0.999`` typical).
epsilonfloat1e-05Small constant for Laplace smoothing of cluster counts.
method

call

Python

Quantize x through all residual levels.

helia_edge/layers/ema_residual_vector_quantizer.py:195

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 of call
NameTypeDefaultDescription
xkeras.KerasTensorRequired``[..., D]`` latent to be quantized.
trainingboolFalseIf ``True``, run EMA codebook updates.
return_indicesboolFalseIf ``True``, also return per-level flat indices.
Returns of call
TypeDescription
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.