Wav2vec-C: A Self-supervised Model for Speech Representation Learning
Samik Sadhu, Di He, Che-Wei Huang, Sri Harish Mallidi, Minhua Wu, Ariya Rastrow, Andreas Stolcke, Jasha Droppo, Roland Maas
Introduction
Self-supervision is a paradigm of machine learning (ML) that deals with unsupervised learning of structural patterns in data by exploiting contextual information. Self-supervision has been of significant interest in the automatic speech recognition (ASR) literature primarily as a pre-training step before a fully supervised task. In particular, it is widely used for problems with some amount of labeled data (for supervised training) and a significantly larger volume of unlabeled data (for self-supervised training). The recently proposed wav2vec 2.0 is one such self-supervised learning model that learns to predict masked out discrete speech encodings using a contextualized representation from a transformer model .
In this paper, we introduce the wav2vec-C model that solves a more rigorously defined self-supervised learning problem compared to the wav2vec 2.0. In the latter, a contrastive loss defined on discretized codes drives the self-supervised learning - including the codebook in the built-in differentiable Vector Quantization module. In contrast, wav2vec-C facilitates codebook learning through an additional regularization on the discrete speech representations by reconstructing the discrete codes to the input features. Thus, wav2vec-C maintains a consistency between the learnt representations and the input features to the network.
We use real world far-field voice query speech with varied degrees of SNR ranging between -40 to 50 dB, whereas most studies on self-supervised learning in the literature use clean read speech and some use simulated noisy speech .
Self-supervised learning has been shown to be useful for settings with little labeled data . It has been observed that the effectiveness of self-supervision decreases as the amount of labeled data increases. In this work, we explore the applicability of self-supervision with a relatively large amount of labeled data (1k hours).
We also limit our model size to facilitate low-latency production level ASR models, which goes against the general trend of exceedingly large self-supervised models proposed in the literature .
We explore and compare different variants of our framework in the choice of the vector quantization framework and the effect it has on robustness and codebook utilization.
The Wav2vec-C Model
Our model is similar to wav2vec 2.0 , but differs in the way we use log short-term Fourier transform (log-STFT) features as input to our model. An encoder network maps the input features to a latent embedding space. These embeddings are quantized by a vector quantization module . The embedded vectors are passed through a SpecAugment module that randomly masks a portion of these embeddings to generate . These masked embeddings are fed into a context network that generates a set of context representation . A contrastive score between the context representations and the vector quantized embeddings is maximized during network training.
2 Wav2vec-C
The wav2vec 2.0 model relies on a diverse set of codes correlating to the underlying speech units learned by to enable to learn good contextual representations via the contrastive loss. However, the wav2vec 2.0 problem formulation can result in several locally optimal codebooks. A few highly probable optima observed in our experiments were
Voice activity detection (VAD) codebooks - where two codes are assigned; one for speech and the other for non-speech
Temporally invariant codebooks - where the model assigns specific codes to fixed temporal locations to enable a good contrastive loss
Our training data consists of many similar query terms occurring at fixed temporal locations which also contributed to the model assigning fixed codes at specific temporal instances via the recurrent encoder (Section 2.3) irrespective of the underlying speech sounds. Hence, the codebook learning methodology adopted for wav2vec 2.0 might not generalize well to other datasets and different model architectures, as in our case.
In wav2vec-C (Figure 1) we enforce the codes to explicitly carry information about the input features to help mitigate the described codebook learning issues. We define an additional consistency network that reconstructs the quantized encodings to consistency vectors and minimize the normed distance between the inputs and during network training. This network allows a flow of information from the input log-STFT features back to the feature domain and enforces the latent space to preserve meaningful information that enable a low reconstruction error. Hence, in a way, wav2vec-C can be seen as an integration of the ideas behind wav2vec 2.0 and VQ-VAE .
3 Encoder network (f𝑓f)
Our encoder network consists of three layers of long short-term memory network (LSTM) with a hidden dimension of 768. The encoder gradients are scaled by a factor as in wav2vec 2.0 to help stabilize the codebook during training.
4 Vector quantization (q𝑞q)
4.2 K-means [12]
During forward pass, a k-means codebook selects the code from which has the closest squared distance to as
However, during back-propagation, a straight-through estimator bypasses gradient computation w.r.t the quantized embedding and copies the gradients to the continuous embedding . Since this process puts the codebook out of the training graph, there are two loss terms incorporated into training as
On minimization of , the first term pushes the quantized representations close to the continuous encoded representation and the second term (also called commitment loss) enforces encodings to commit to quantized embeddings during training. In eq. 2, is the stop gradient operator and as is the optimal value reported in .
5 Masking
We use a SpecAugment module to mask out portions of the continuous encodings before feeding them to the context network. We use five masks for every utterance. Each mask has maximum width of 16% of the utterance length. On average 40% of the encoded frames are masked.
6 Context network (g𝑔g)
The context network consists of five transformer layers, with model dimension 1024 and inner feed-forward dimension of 4096 with 16 attention heads. We use sinusoidal positional embedding for the transformer layers. A contrastive score between the context representations and the quantized encodings is computed as
where , is a set consisting of and a selection of negative samples, is the temperature variable and calculates the cosine similarity . In our experiments, we uniformly sample negative samples from the encodings of the utterance and is updated as proposed in .
7 Consistency network (r𝑟r)
The consistency network consists of a 3-layer LSTM that maps the quantized embedding to the consistency vectors . We minimize the normed distance between and as
8 Loss
During training, we minimize the primary contrastive loss together with a codebook loss component and the consistency loss as
The codebook loss (section 2.4) takes a different form according to the type of VQ used. Wav2vec 2.0 and wav2vec-C are generalized by the parameter , where a value results in the wav2vec 2.0 model as the consistency loss is ignored for model training, while leads to our wav2vec-C model in full effect.
For a Gumbel-softmax VQ module, the codebook loss is given by , where is a diversity loss on the Gumbel-softmax distribution given by
where is the probability assignment by the codebook on the code. The weight on the diversity loss determines the relative importance of the component and is instrumental in avoiding the codebook collapse that is commonly observed in VQ problems . In our experiments, we found to be suitable to avoid catastrophic codebook collapse issues. For k-means VQ, the codebook loss is simply equal to the k-means loss, i.e.,
Experimental Setup
The goal of this study is to evaluate the effectiveness of self-supervised pre-training for real world applications. Hence, instead of using publicly available clean read speech we use in-house training and evaluation data consisting of real-world far-field English voice command and voice query speech collected from home environments similar to with varying degrees of SNR in the range -40 to 50 dB.
We use 10k hours of unlabeled and 1k hours of transcribed de-identified English language training data collected from native and non-native English speakers. To our knowledge, this work is one of the first few instances where a large proportion of labeled data is used alongside self-supervised pre-training for ASR tasks, especially realistic speech queries instead of clean read speech data.
1.2 Test data
We test our ASR models on four different test sets summarized in Table 1
2 Recurrent Neural Network Transducer (RNN-T) Model
RNN-T ASR models are widely used for deployable end-to-end speech recognition systems because of their fast online streaming capability. We use the pre-trained wav2vec-C and wav2vec 2.0 models to initialize the speech encoder for a RNN-T ASR model.
After training the self-supervised model on unlabeled data, we use the output of the context network as speech representations. Thus, the RNN-T speech encoder consists of three LSTM layers followed by five layers of transformer extracted from the self-supervised model with the masking module eliminated. We use two LSTM layers with 1024 hidden units as the RNN-T prediction network and a simple single layer feedforward joint network. The pre-trained speech encoder is also fine-tuned during RNN-T training. We use a total of 4000 sub-word tokens together with a blank token to generate the targets for RNN-T training. The RNN-T network is also regularized with SpecAugment on the input features with 10% of the temporal frames and 30% of the frequency bins randomly masked with noise. 25% dropout is applied on the transformer weights.
2.2 Baseline RNN-T
Our baseline model consists of an RNN-T with the same architecture as the pre-trained model but without pre-training the speech encoder.
3 Training details
We train 4 different self-supervised models
wav2vec 2.0 (GS): , Gumbel-softmax codebook
wav2vec 2.0 (KM): , k-means codebook
wav2vec-C (GS): , Gumbel-softmax codebook
wav2vec-C (KM): , k-means codebook
Subsequently, we train RNN-T models with the speech encoder replaced by the self-supervised models.
Our models are trained using Tensorflow 2.0. The self-supervised models are trained for 100k steps with 30 minutes of speech per step. We use an Adam optimizer , where the learning rate is warmed up from and held at after 3k steps. The RNN-T models are trained for 60k steps, with an average of 1 hours of speech per step. The learning rate is warmed up from and held at after 3k steps.
Results
We compare the word error rate reduction relative to the baseline model (rWERR) for the different pre-trained RNN-T models evaluated on the four test sets in Table 2. The baseline ASR model has absolute word error rate. To smooth out error fluctuations, we report the mean rWERR computed after 50k, 55k and 60k RNN-T training steps. The average rWERR in the last column is the rWERR for each test set weighted by the number of utterances in that test set.
Our implementation of the wav2vec 2.0 pre-trained RNN-T model does not show noticeable performance improvement over baseline for the clean test sets. Whereas, for the noisy test sets, some gains can be observed - with wav2vec 2.0 (KM) performing better, on average, compared to wav2vec 2.0 (GS). This trend is comparable to the results reported in , where pre-training is shown to be most beneficial for the challenging test_other test set of Librispeech . However, while drawing this comparison we should keep in mind the major differences between the best performing wav2vec 2.0 models in and our implementation, namely
We use a much smaller context network (5 layers) compared to the original (24 layers)
We use a 3-layer LSTM as encoder with log-STFT input features
The wav2vec-C encoded RNN-T models, on the other hand, show a positive rWERR for both as well as clean test sets. In particular, wav2vec-C (GS) gains 1.6% rWERR on and 1.2% rWERR on . However, there is a reduction in performance (in comparison to wav2vec 2.0) for the noisy test sets. This suggests that the reconstruction idea adopted for wav2vec-C leads to an overall better performance of the pre-trained RNN-T model, however with a slight loss in robustness.
Our codebooks have a maximum capacity of k codes with wav2vec-C (GS) utilizing the full 100% of the codebook (see Table 3). Hence, the consistency loss together with the weight on the diversity loss enforces the model to pick a variety of codes to minimize the reconstruction loss.
A t-SNE plot of the 102.4k codes in the 100% utilized codebook of the wav2vec-C (GS) model can be seen in Figure 2(b) showing the clusters formed by the codes over the course of training. On the other hand, the 102.4k codes learnt by wav2vec 2.0 (GS), as shown in Figure 2(a), form a smaller number of clusters with significant inter-cluster overlap possibly due to the under-utilized codebook.
The k-means codebook uses only a small fraction of the codes but is more robust compared to Gumbel-softmax models for noisy test sets, in particular the noisy test set. For example, a comparison of the ASR performances of wav2vec-C (GS) and wav2vec-C (KM) would show that wav2vec-C (GS) gives a better rWERR for clean test sets in comparison to noisy test sets, whereas wav2vec-C (KM) shows the opposite characteristics. This observation highlights the importance of codebook diversity for different application domains. For example, a small codebook diversity is not necessarily a bad design choice if robustness is of importance during model evaluation.
Conclusions
In this paper we propose wav2vec-C, a new self-supervised learning model which is based on an amalgamation of the ideas from wav2vec 2.0 and VQ-VAE with the goal of solving the codebook utilization difficulties observed for wav2vec 2.0. We used real-world far-field noisy data for self-supervised learning and 1k hours of data for supervised ASR training. The proposed self-supervised model after RNN-T fine-tuning achieved, on average, a 1.4% relative WER reduction over baseline compared to a 0.7% reduction from wav2vec 2.0. Furthermore, we also observed that ASR robustness is correlated with codebook diversity, validating our motivation for the wav2vec-C architecture