espnet2.tok.quantizer.differentiable_kmeans.DifferentiableKMeans
espnet2.tok.quantizer.differentiable_kmeans.DifferentiableKMeans
class espnet2.tok.quantizer.differentiable_kmeans.DifferentiableKMeans(centroid_path: str | Path, distance_type: str = 'squared_euclidean', sigma_squared: float = 1.0, temperature_init: float = 2.0, temperature_floor: float = 0.1, temperature_decay: float = 0.999995)
Bases: AbsSpeechTokenizerQuantizer
Quantize SSL features using trainable k-means centroids.
Training uses Gumbel-Softmax straight-through assignments whose forward values are hard one-hot vectors. The soft backward path permits joint optimization of cluster centroids and the upstream SSL frontend. Its results are exposed through the shared speech-tokenizer output type.
Initialize trainable centroids from an offline k-means model.
- Parameters:
- centroid_path – Joblib model produced by ESPnet’s k-means pipeline. The loaded object must expose
cluster_centers_with shape(num_clusters, feature_dim). - distance_type – Distance used to construct assignment logits. Either
"squared_euclidean"or"euclidean". - sigma_squared – Scale applied to negative distance logits.
- temperature_init – Initial Gumbel-Softmax temperature.
- temperature_floor – Minimum annealed temperature.
- temperature_decay – Exponential decay applied per training forward.
- centroid_path – Joblib model produced by ESPnet’s k-means pipeline. The loaded object must expose
encode(features: Tensor, feature_lengths: Tensor) → SpeechTokenizerOutput
Apply deterministic nearest-centroid assignment without Gumbel noise.
property feature_dim : int
Return the centroid feature dimension.
forward(features: Tensor, feature_lengths: Tensor) → SpeechTokenizerOutput
Sample hard straight-through assignments for differentiable training.
property num_clusters : int
Return the number of trainable centroids.
set_temperature(temperature: float) → None
Set the current temperature without changing annealing parameters.
