Effectiveness of self-supervised pre-training for speech recognition

Alexei Baevski, Michael Auli, Abdelrahman Mohamed

Introduction

Representation learning has been an active research area for more than 30 years , with the goal of learning high level representations which separates different explanatory factors of the phenomena represented by the input data . Building Automatic Speech Recognition (ASR) systems, typically requires a large volume of training data to represent different factors contributing to the creation of speech signals, e.g. background noise, recording channel, speaker identity, accent, emotional state, topic under discussion, and the language used in communication. The practical need for building ASR systems for new conditions with limited resources spurred a lot of work focused on unsupervised speech recognition and representation learning , in addition to semi- and weakly-supervised learning techniques to reduce the supervised data needed in real-world scenarios .

Recently impressive results have been reported for representation learning, that generalizes to different downstream tasks, through self-supervised learning for NLP and speech . Self-supervised representation learning tasks include predicting masked parts of the input, reconstructing inputs through low bit-rate channels, or contrasting similar data points against different ones.

In this work we compare different approaches of self-supervised pre-training for speech data. We consider learning discrete units to represent the audio data through either self-supervision or through clustering spectral features, followed by pre-training over these units using a bi-directional transformer (BERT) model . This is compared to directly learning representations without explicit quantization over the raw audio as well as spectral features. Previous work fed the representations of the pre-trained model into a task-specific architecture for speech recognition instead of the raw waveform . Instead we directly fine-tune the pre-trained BERT model on transcribed speech data using a CTC loss. Our experiments demonstrate that discrete unit discovery, followed by BERT training achieves better results than representations learned without explicit quantization. Disentangling acoustic unit discovery from learning the sequential relationship between them, enables better representations of the data which in turn improves down-stream model accuracy.

We pre-train our models on the unlabeled 960h Librispeech data and follow the Libri-light limited resource supervised training sets of 10 hours, 1 hour, and 10 mins. Our best model fine-tuned only on 1 hour of labeled data can outperform the best known result from the literature relying on 100h of labeled data on the standard Librispeech test-other subset. Using only 10 minutes of labeled data the approach achieves 16.3/25.2 WER on test-clean/other.

Preliminaries

Using self-supervision, BERT , a deep bidirectional transformer model, builds its internal language representation that generalizes to other downstream NLP tasks. Self-attention over the whole input word sequence enables BERT to jointly condition on both the left and right context of data. For training, it uses both a masked language model loss, by randomly removing some input words for the model to predict, and a contrastive loss to distinguish the next sentence in the document from a randomly selected one.

2 Wav2Vec

where TT is the sequence length, σ(x)=1/(1+exp⁡(−x))\sigma(x)=1/(1+\exp(-x)), and where σ(zi+k⊤hk(ci))\sigma(\mathbf{z}_{i+k}^{\top}h_{k}(\mathbf{c}_{i})) is the probability of zi+k\mathbf{z}_{i+k} being the true sample. A step-specific affine transformation hk(ci)=Wkci+bkh_{k}(\mathbf{c}_{i})=W_{k}\mathbf{c}_{i}+\mathbf{b}_{k} is applied to ci\mathbf{c}_{i} . The loss L=∑k=1KLk\mathcal{L}=\sum_{k=1}^{K}\mathcal{L}_{k} is optimized by summing (1) over different step sizes. The learned high level features produced by the context network ci\mathbf{c}_{i} are shown to be better acoustic representations for speech recognition compared to standard spectral features.

3 vq-wav2vec

vq-wav2vec learns vector quantized (VQ) representations of audio data using a future time-step prediction task. Similar to Wav2Vec, there are a convolutional encoder and decoder networks f:X↦Zf:\mathcal{X}\mapsto\mathcal{Z} and g:Z^↦Cg:\hat{\mathcal{Z}}\mapsto\mathcal{C} for feature extraction and aggregation. However, in between them there is a quantization module q:Z↦Z^q:\mathcal{Z}\mapsto\hat{\mathcal{Z}} to build discrete representations which are input to the aggregator.

Approach

Our work builds on the recently proposed work in where audio is quantized using a contrastive loss, then features learned on top by a BERT model . For the vq-wav2vec quantization, we use the gumbel-softmax variant with the same setup as described in . This model quantizes the Librispeech dataset into 13.5k unique codes.

To understand the impact of discrete acoustic representations of vq-wav2vec , as alternatives, we explore quantizing the standard mel-frequency cepstral coefficients (MFCC) and log-mel filterbanks coefficients (FBANK), choosing a subset small enough to fit into GPU memory and running k-means with 13.5k centroids (to match the vq-wav2vec setup) to convergence. We then assign the index of the closest centroid to represent each time-step.

We train a standard BERT model with only the masked language modeling task on each set of inputs in a similar way as described in , namely by choosing tokens for masking with probability of 0.05, expanding each chosen token to a span of a length sampled from a normal distribution with mean 10 and standard deviation 10 (spans may overlap) and then computing a cross-entropy loss which attempts to maximize the likelihood of predicting the true token for each one that was masked (Figure 1(a)).

Following , we replace the fixed positional embeddings in the BERT model with a single group convolutional layer that is applied directly on the embeddings before any of the transformer blocks. The convolutional layer has a kernel size of 128 and group size of 16 to reduce the number of added parameters.

2 Continuous BERT

where each sample zi\mathbf{z}_{i} is computed as a dot product of the output of the model at timestep ii and the true unmasked value of positive example at timestep ii or a randomly sampled negative example. To stabilize training, we add the squared sum of logits produced by the dot-product to the loss, and then apply a soft clamp si^=λtanh⁡(si/λ)\hat{s_{i}}=\lambda\tanh(s_{i}/\lambda) for each logit sis_{i} to prevent the model’s tendency to continually increase the magnitude of logits during training . We use the same kind of convolutional positional layer as described in section 3.1.

3 Supervised fine-tuning

The pre-trained models are fine-tuned to perform the ASR task by adding a randomly initialized linear projection on top of the features computed by the transformer models into VV classes representing the vocabulary of the task. The vocabulary is 29 tokens for character targets plus a word boundary token. The models are optimized by minimizing the CTC loss. We apply SpecAugment inspired masking to time-steps and channels during training which delays overfitting and significantly improves the final accuracy numbers, especially on the smallest subsets.

We train a single seed for all subsets except the 10 minute one, where we train 5 seeds and choose the best one. For 10 minute subset, some of the seeds fail by entering the overfit regime very early. We train on a single GPU using the Adam optimizer with a tri-state learning scheduler where the learning rate is linearly increased from 1e-7 to 2e-05 in the first stage, held at 2e-5 in the second, and finally linearly decayed to 0 in the third. For 1h and 10min subsets we train for 1250 / 6600 / 12150 updates in each respective stage, for 10h subset we train for 5000 / 16500 / 28500 updates and for 100h we train for 8000 / 52800 / 91200 updates. We use a batch size of 6144 timesteps (61.44 seconds worth of audio) for the 100 hour subset and 3072 timesteps for other subsets.

During fine-tuning, we apply a modified SpecAugment policy, where we randomly choose a number of starting timesteps to mask, with each timestep having a chance of 3.75% of being chosen. A span of 20 timesteps starting at each of the chosen position is then replaced with the mask embedding used during unsupervised training (spans may overlap). We also apply channel masking, in which we choose the starting channel index from all channels with a probability of 0.4% and the length of the channel mask by sampling from a normal distribution with a mean of 64 and standard deviation of 64. The chosen (and possibly overlapping) spans of channels are then zeroed out.

We apply a dropout at every layer of the transformer of 0.1 for 10 minute and 1 hour setup, but we disable it for other subsets as masking described above appears to provide enough regularization.

Experiments

We implement our models in the fairseq toolkit.

All experiments are performed by pre-training on the 960 hours of audio only data of the Librispeech training set, fine-tuning on the Libri-light limited resource supervised training sets of 10 hours (24 speakers), 1 hour (24 speakers), and 10 minutes (4 speakers). The Libri-light training sets are sampled equally from the two clean and noisy portions, a balance of male and female speakers. We also report results of models fine-tuned on 100 hours following the “train-clean-100” subset. All models are evaluated on the standard Librispeech dev and test splits.

2 Models

We first train the vq-wav2vec quantization model following the gumbel-softmax setup described in . After training this model on 960h of Librispeech and quantizing the training dataset, we are left with 13.5k unique codewords combinations.

For quantizing MFCC and FBANK features extracted using the Kaldi toolkit, we use 8 Volta GPUs with 32GB memory to compute 13.5k K-Means centroids matching the number of unique tokens produced by the vq-wav2vec model. To fit into GPU memory, we subsample 50% of MFCC features and 25% of FBANK features from the training set before running the clustering algorithm.

The model we use for the masked language modeling task is a standard BERT model with 12 transformer layers, model dimension 768, inner dimension (FFN) 3072 and 12 attention heads . The learning rate is warmed up over the first 10,000 updates to a peak value of 1\times10−51\text{\times}{10}^{-5}, and then linearly decayed over a total of 250k updates. We train on 128 GPUs with a batch size of 3072 tokens per GPU giving a total batch size of 393k tokens where each token represents 10ms of audio data.

To mask the input sequence, we follow and randomly sample p=0.05p=0.05 of all tokens to be a starting index, without replacement, and mask MM consecutive tokens from every sampled index; spans may overlap. MM is sampled from a Gaussian distribution with μ=10\mu=10 and σ=10\sigma=10, rounded to the nearest integer greater than or equal to zero.

Different from , we do not concatenate different utterances to form examples for training, instead each utterance is treated as a single example, as we find that this approach produces better results after fine-tuning.

2.2 Continuous Inputs Training

For training on dense features, we use a model similar to a standard BERT model with the same parameterization as the one used for quantized input training, but we use the wav2vec, MFCC or FBANK inputs directly. We add 128 relative positional embeddings at every multi-head attention block instead of fixed positional embeddings to make it easier to handle longer examples. We train this model on 8 GPUs with a batch size of 9,600 inputs per GPU, resulting in a total batch size of 76,800. We find that increasing the number of GPUs (which increases the effective batch size) does not lead to better results with this particular setup.

Wav2vec features are 512-dimensional, while MFCC features have 39 dimensions and FBANK features have 80. We introduce a simple linear projection from the feature dimension to BERT dimension (768) for all models.

Similar to 4.2.1, we mask time-steps by randomly sampling, without replacement, p=0.05p=0.05 of all time-steps to be a starting index, and mask MM consecutive time-steps from every sampled index; spans may overlap, where MM is sampled from a Gaussian distribution with μ=10\mu=10 and σ=10\sigma=10, rounded to the nearest integer greater than or equal to zero. We sample 1010 negative examples from other masked time-steps from the same example, and an additional 1010 negative examples from masked time-steps occurring anywhere in the batch. We compute a dot product between the original features and the output corresponding to the same time-step after they are processed by the BERT model. We add the squared sum of logits from these computations multiplied by λ=0.04\lambda=0.04 to the loss, and then apply a smooth clamp by recomputing each logit si^=20tanh⁡(si/20)\hat{s_{i}}=20\tanh(s_{i}/20).

The learning rate is warmed up over the first 10,000 updates to a peak value of 1\times10−51\text{\times}{10}^{-5}, and then linearly decayed over a total of 250k updates.

3 Methodology

For quantized inputs, we compute token indices using the gumbel-softmax based vq-wav2vec model. For MFCC and FBANK features we take the index of the closest centroid, as measured by finding the minimum Euclidean distance, to each corresponding feature in the Librispeech dataset. We then train a BERT model as descirbed in §4.2.1.

For wav2vec continuous inputs, we use features extracted by the publicly available wav2vec model which contains 6 convolutional blocks in the feature extractor and 11 convolutional blocks in the aggregator module. We use the outputs of the aggregator as features. For MFCC and FBANK, we use those features directly after applying a single linear projection to upsample them to the model dimensionality.

We fine-tune our pre-trained models on either 100 hours of Librispeech train-clean-100 subset, 10 hours, 1 hour, or 10 minutes of labelled data following the Libri-light limited training sets. We use the standard CTC loss and train for up to 20k updates. We find that the pre-trained models converge after only around 4k updates, while the models trained from scratch tend to converge much later, around 18k updates. We fine-tune all models with a learning rate of 0.00010.0001 that is linearly warmed up over the first 2k updates and then annealed following a cosine learning rate schedule over the last 18k updates. We set the dropout of the pre-trained BERT models to 0.1 and sweep on dropout of the BERT model outputs before the final projection layer over values between 0.0 and 0.4 in increments of 0.1. For each model, we choose a single best checkpoint that has the best loss on the validation set, which is a combination of dev-clean and dev-other standard Librispeech splits.

We use the publicly available wav2letter++ decoder integrated into the Fairseq framework with the official Librispeech 4-gram language model. We run a sweep on weights for language model score, word score and silence token weights for each model, where parameters are chosen randomly and evaluated on the dev-other Librispeech set. We use the weights found by these sweeps to evaluate and report results for all other splits. The sweeps are run with beam size of 250, while the final decoding uses a beam size of 1500.

4 Results

In our first experiment, we compare unit discovery followed by BERT training over the resulting discrete units (Discrete BERT) to directly learning representations from the audio inputs (Continuous BERT) in different simulated labeled data scenarios ranging from 100 hours to 10 minutes. We compare quantization with vq-wav2vec to clustered MFCC and FBANK features. The continuous BERT variant learns directly from the audio representations without explicit quantization and we experiment with inputting wav2vec, MFCC and FBANK features.

Table 1 compares WERs of different input features and pre-training methods on the standard Librispeech clean and other subsets. The first observation is that Discrete BERT outperforms Continuous BERT in all settings. This shows that pre-training over meaningful discrete units outperforms directly learning representations from the continuous unlabeled data. The initial unit discovery builds a vocabulary that makes the subsequent BERT pre-training more effective.

The best input features are obtained through self-supervised learning through vq-wav2vec for Discrete BERT, or wav2vec for Continuous BERT.

For Discrete BERT, vq-wav2vec provides about 40% of relative error reduction for both test sutsets compared to clustered spectral features across all training set sizes, with bigger gains on the noisy test-other subset.

Pre-training brings clear benefits: When reducing the amount of labeled training data from 100h to 10h results in an increase of only 2 WER on test-other and 1.4 WER on test-clean for Discrete BERT with vq-wav2vec inputs. This shows that pre-training is effective and particularly so when little labeled data is available. When reducing the amount of labeled data to only 10 minutes, Discrete BERT with va-wav2vec inputs can still achieve a WER of 16.3/25.2 on test-clean/other.

Table 2 shows a comparison of Discrete BERT to results from the literature. Fine-tuning Discrete BERT on only 10 hour of labeled data can nearly match the best known result on 100 hours of labeled Librispeech data on test-clean, while achieving a 25% relative WER reduction on test-other. Moreover, when using the same train-clean-100 subset for fine-tuning, Discrete BERT with vq-wav2vec inputs improves by 6.5 WER (35% relative WER reduction) on test-other and 1.3 WER (22% relative WER reduction) on test-clean over .

The closest setup to ours is which learns representations using CPC and then feed these into an ASR system trained on about 96h of labeled data, which is surpassed on test-clean by our Discrete-BERT model trained on vq-wav2vec representations fine-tuned on 10mins and 1 hour.

5 Ablations

To better understand the impact of BERT pre-training in our representation learning approach, we remove the BERT pre-training step and only perform unit discovery through vq-wav2vec, for discrete inputs, and fine-tuning, for both discrete and continuous inputs on the 10 hour labeled setup. The vocabulary for discrete inputs is still built on the unlabeled data. Table 3 shows that training with discrete inputs fails. This is likely because the representations of the input discrete units are random and training on the labeled data is not sufficient. Continuous inputs do not suffer from this issue.

Next, we shed some light on how a two-step pre-training approach compares to a single-step approach. Specifically, we compare Continuous BERT with wav2vec input features (requiring separate learning of the wav2vec features) to just wav2vec features fine-tuned with a CTC loss on labeled data. The results (Table 4) show that Continuous BERT + wav2vec provides substantial gains. A second step of representation learning more than halved the WER, with more gains observed in the “clean” subset (cf. 4.4).

Discussion and Related Work

The success of BERT and Word2Vec for NLP tasks motivated more research on self-supervised approaches for acoustic word embedding and unsupervised acoustic feature representation , either by predicting masked discrete or continuous input, or by contrastive prediction of neighboring or similarly sounding segments using distant supervision or proximity in the audio signal as an indication of similarity. In a dynamic time warping alignment is used to discover similar segment pairs.

Our work is inspired by research efforts reducing the dependence on labeled data for building ASR systems through unsupervised unit discovery and acoustic representation leaning , and through multi- and cross-lingual transfer learning in low-resource conditions , and semi-supervised learning .

Conclusion and Future work

We presented a systematic comparison of self-supervised pre-training approaches for speech recognition. The most effective method is to first learn a discrete vocabulary of the data with vq-wav2vec followed by standard BERT training over these discrete units. This performs much better than directly learning from the continuous audio data. Different to previous work which relied on task-specific ASR models, we directly fine-tune the resulting BERT model on transcribed speech data to act as speech recognition models. This approach can achieve better accuracy on test-other than the best known result with 100 hours of labeled data while relying on two orders magnitude less labeled data. When the model is fine-tuned on only 10 minutes of data, it can still achieve a WER 25.2 on test-other and WER 16.3 on test-clean.

References