class
GumbelSoftmaxBottleneck
PythonDiscrete bottleneck via Gumbel-Softmax (Concrete) with optional straight-through hard one-hot.
GumbelSoftmaxBottleneck( num_embeddings: int, embedding_dim: int, temperature: float = 1.0, hard: bool = True, input_is_logits: bool = False, use_bias: bool = True, kl_weight: float = 1.0, **kwargs={},)Discrete bottleneck via Gumbel-Softmax (Concrete) with optional straight-through hard one-hot.
Tracks metrics (logged via metrics):
- gs_bits_per_index (lower bound, bits/index)
- gs_perplexity (empirical perplexity from hard argmax histogram)
- gs_usage (fraction of codes used at least once in the batch)
- gs_temperature (current τ; useful when annealing)
constant
K
PythonK = int(num_embeddings)constant
D
PythonD = int(embedding_dim)attribute
hard
Pythonhard = bool(hard)attribute
input_is_logits
Pythoninput_is_logits = bool(input_is_logits)attribute
use_bias
Pythonuse_bias = bool(use_bias)attribute
kl_weight
Pythonkl_weight = float(kl_weight)attribute
tau
Pythontau = self.add_weight(name='temperature', shape=(), initializer=keras.initializers.Constant(float(temperature)), trainable=False, dtype='float32')attribute
metrics
Pythonmetricsmethod
build
Pythonbuild(input_shape)method
set_temperature
Pythonset_temperature(value: float)method
call
Pythoncall(x, training=False, return_indices: bool = False, return_probs: bool = False)method
get_config
Pythonget_config()