VQAutoencoder
PythonConvenience wrapper around (encoder -> VectorQuantizer -> decoder).
VQAutoencoder(encoder: keras.Model, vq: VectorQuantizer, decoder: keras.Model, **kwargs={})Convenience wrapper around (encoder -> VectorQuantizer -> decoder).
- Supports extra reconstruction-side losses and metrics.
- Can return discrete code indices from the VQ bottleneck.
- Exposes VQ layer metrics alongside base model metrics.
Parameters
| Name | Type | Default | Description |
|---|---|---|---|
encoder | keras.Model | Required | Encoder model producing continuous latents. |
vq | VectorQuantizer | Required | VectorQuantizer layer that discretizes latents. |
decoder | keras.Model | Required | Decoder model mapping bottleneck outputs to reconstructions. |
encoder
Pythonencoder = encodervq
Pythonvq = vqdecoder
Pythondecoder = decodermetrics
Pythonmetricscall
PythonRun encoder -> VQ bottleneck -> decoder.
call( x: keras.KerasTensor, training: bool = False, return_indices: bool = False,) -> keras.KerasTensor | tuple[keras.KerasTensor, keras.KerasTensor]Run encoder -> VQ bottleneck -> decoder.
Parameters
| Name | Type | Default | Description |
|---|---|---|---|
x | keras.KerasTensor | Required | Input batch. |
training | bool | False | Whether to run in training mode (affects encoder/decoder/VQ). |
return_indices | bool | False | If True, also return discrete code indices. |
Returns
| Type | Description |
|---|---|
keras.KerasTensor | tuple[keras.KerasTensor, keras.KerasTensor] | Reconstruction, optionally with indices. |
compile
PythonCompile with optional extra losses/metrics.
compile( optimizer: keras.optimizers.Optimizer, loss: keras.losses.Loss | None = None, metrics: list | None = None, extra_losses: list | None = None, extra_metrics: list | None = None, **kwargs={},)Compile with optional extra losses/metrics.
Parameters
| Name | Type | Default | Description |
|---|---|---|---|
optimizer | keras.optimizers.Optimizer | Required | Keras optimizer |
loss | keras.losses.Loss | None | None | base reconstruction loss (e.g., keras.losses.MeanSquaredError()) |
metrics | list | None | None | standard Keras metrics (Metric instances or callables) |
extra_losses | list | None | None | list[Callable(y_true, y_pred) -> scalar] |
extra_metrics | list | None | None | list of Metric OR Callable(y_true, y_pred) -> scalar |
compute_loss
PythonCompute total loss = recon + extra losses + layer-added losses.
compute_loss(x=None, y=None, y_pred=None, sample_weight=None, allow_empty=False)Compute total loss = recon + extra losses + layer-added losses.
get_config
PythonReturn serialized config for saving/loading.
get_config()Return serialized config for saving/loading.
from_config
PythonRecreate model from serialized config.
from_config(config, custom_objects=None)classmethod
Recreate model from serialized config.
compute_metrics
PythonUpdate compiled metrics plus extra metric trackers.
compute_metrics(x, y, y_pred, sample_weight=None)Update compiled metrics plus extra metric trackers.