GSAutoencoder
PythonConvenience wrapper around (encoder -> GumbelSoftmaxBottleneck -> decoder).
GSAutoencoder(encoder: keras.Model, gs: GumbelSoftmaxBottleneck, decoder: keras.Model, **kwargs={})Convenience wrapper around (encoder -> GumbelSoftmaxBottleneck -> decoder).
- Supports extra reconstruction-side losses and metrics.
- Can return discrete code indices and/or code probabilities from the bottleneck.
- Exposes Gumbel-Softmax layer metrics alongside base model metrics.
Parameters
| Name | Type | Default | Description |
|---|---|---|---|
encoder | keras.Model | Required | Encoder model producing continuous latents. |
gs | GumbelSoftmaxBottleneck | Required | GumbelSoftmaxBottleneck layer that discretizes latents. |
decoder | keras.Model | Required | Decoder model mapping bottleneck outputs to reconstructions. |
encoder
Pythonencoder = encodergs
Pythongs = gsdecoder
Pythondecoder = decodermetrics
Pythonmetricscall
PythonRun encoder -> GS bottleneck -> decoder.
call( x: keras.KerasTensor, training: bool = False, return_indices: bool = False, return_probs: bool = False,) -> keras.KerasTensor | tuple[keras.KerasTensor, ...]Run encoder -> GS 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/gs). |
return_indices | bool | False | If True, also return the discrete code indices. |
return_probs | bool | False | If True, also return code probabilities. |
Returns
| Type | Description |
|---|---|
keras.KerasTensor | tuple[keras.KerasTensor, ...] | Reconstruction, optionally with indices and/or probabilities. |
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 of callables (y_true, y_pred) -> scalar to add to loss. |
extra_metrics | list | None | None | Metric instances or callables (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.
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.
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.