Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
180 changes: 180 additions & 0 deletions recipes/LibriSpeech/ASR/CTC/extracted_features/extract_features.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,180 @@
import speechbrain as sb
from speechbrain.utils.logger import get_logger
from speechbrain.dataio.dataloader import LoopedLoader
from speechbrain.dataio.dataloader import DataLoader
import torch
import numpy as np
from tqdm import tqdm
from abc import ABCMeta, abstractmethod
from pathlib import Path
from contextlib import ExitStack
from typing import Dict, Any
import os
import lilcom
from dataclasses import dataclass
from speechbrain.dataio.feature_io import FeatureStorageConfig, FeatureStorageWriter, create_feature_storage_writers, NumpyHdf5Writer
from hyperpyyaml import load_hyperpyyaml
import sys
"""
python extract_features.py hparams/extract_ssl_representations.yaml --data_folder=$SLURM_TMPDIR/librispeech/LibriSpeech/ --output_folder $SLURM_TMPDIR/results/extract_ssl_representations/ --batch_size=64


python train_with_wav2vec.py hparams/train_hf_wav2vec.yaml --data_folder=$SLURM_TMPDIR/LibriSpeech/ --output_folder $SCRATCH/results/wav2vec2-base-960h/ --extracted_features_folder $SLURM_TMPDIR/results/extract_ssl_representations/ssl_features --batch_size=32

find . -type f -name '*.tar.gz' -exec tar -xzf {} -C . \;
scp -r $HOME/projects/def-ravanelm/datasets/librispeech/* .


python extract_features.py hparams/extract_ssl_representations.yaml --data_folder=$SLURM_TMPDIR/LibriSpeech/ --output_folder $SCRATCH/extracted_features/wav2vec2-base-960h/ --batch_size=20
python extract_features.py hparams/extract_ssl_representations.yaml --data_folder=$SLURM_TMPDIR/LibriSpeech/ --output_folder $SCRATCH/extracted_features/hubert-large-ll60k/ --batch_size=20

python train_with_wav2vec.py hparams/train_hf_wav2vec.yaml --data_folder=$SLURM_TMPDIR/LibriSpeech/ --output_folder $SCRATCH/results/wav2vec2-base-960h-10-epochs/ --extracted_features_folder $SLURM_TMPDIR/results/extract_ssl_representations/save/ssl_features --batch_size=32 --number_of_epochs 10
"""
logger = get_logger(__name__)

@dataclass
class FeatureExtractionConfig:
utterance_id_key: str = "id"
ssl_key: str = "ssl_feats"

class ExtractFeatures(sb.core.Brain):
def __init__(
self,
modules,
hparams,
run_opts,
feature_extraction_config: FeatureExtractionConfig):
super().__init__(
modules=modules,
hparams=hparams,
run_opts=run_opts,
)
self.feature_extraction_config = feature_extraction_config

def compute_features(self, batch, stage):
batch = batch.to(self.device)
wavs, wav_lens = batch.sig
batch_size = wavs.shape[0]

# extract features
feats = self.modules.wav2vec2(wavs, wav_lens)

return [
{
self.feature_extraction_config.utterance_id_key: batch.id[i],
self.feature_extraction_config.ssl_key: feats[i],
} for i in range(batch_size)
]

def dataio_prepare(hparams):
"""This function prepares the datasets to be used in the brain class.
It also defines the data processing pipeline through user-defined functions.
"""
data_folder = hparams["data_folder"]

train_data = sb.dataio.dataset.DynamicItemDataset.from_csv(
csv_path=hparams["train_csv"],
replacements={"data_root": data_folder},
)
valid_data = sb.dataio.dataset.DynamicItemDataset.from_csv(
csv_path=hparams["valid_csv"],
replacements={"data_root": data_folder},
)
# test is separate
test_datasets = {}
for csv_file in hparams["test_csv"]:
name = Path(csv_file).stem
test_datasets[name] = sb.dataio.dataset.DynamicItemDataset.from_csv(
csv_path=csv_file, replacements={"data_root": data_folder}
)

datasets = [train_data, valid_data] + [i for k, i in test_datasets.items()]

# 2. Define audio pipeline:
@sb.utils.data_pipeline.takes("wav")
@sb.utils.data_pipeline.provides("sig")
def audio_pipeline(wav):
sig = sb.dataio.dataio.read_audio(wav)
return sig

sb.dataio.dataset.add_dynamic_item(datasets, audio_pipeline)
# 4. Set output:
sb.dataio.dataset.set_output_keys(
datasets,
["id", "sig"],
)

return train_data, valid_data, test_datasets

if __name__ == "__main__":
# CLI:
hparams_file, run_opts, overrides = sb.parse_arguments(sys.argv[1:])

with open(hparams_file, encoding="utf-8") as fin:
hparams = load_hyperpyyaml(fin, overrides)

# Create experiment directory
sb.create_experiment_directory(
experiment_directory=hparams["output_folder"],
hyperparams_to_save=hparams_file,
overrides=overrides,
)

# Dataset prep (parsing Librispeech)
from librispeech_prepare import prepare_librispeech

sb.utils.distributed.run_on_main(
prepare_librispeech,
kwargs={
"data_folder": hparams["data_folder"],
"tr_splits": hparams["train_splits"],
"dev_splits": hparams["dev_splits"],
"te_splits": hparams["test_splits"],
"save_folder": hparams["output_folder"],
"merge_lst": hparams["train_splits"],
"merge_name": "train.csv",
"skip_prep": hparams["skip_prep"],
},
)

# here we create the datasets objects as well as tokenization and encoding
train_data, valid_data, test_datasets = dataio_prepare(
hparams
)

feature_extractor = ExtractFeatures(
modules=hparams["modules"],
hparams=hparams,
run_opts=run_opts,
feature_extraction_config=FeatureExtractionConfig(
utterance_id_key="id",
ssl_key="ssl_feats",
)
)

feature_extractor.cache_features(
hparams["train_feature_storage_writers"],
train_data,
loader_kwargs=hparams["dataloader_opts"],
stage=sb.Stage.TRAIN
)
feature_extractor.cache_features(
hparams["valid_feature_storage_writers"],
valid_data,
loader_kwargs=hparams["dataloader_opts"],
stage=sb.Stage.VALID
)
for k, v in test_datasets.items():
feature_extractor.cache_features(
hparams["test_feature_storage_writers"][k],
v,
loader_kwargs=hparams["dataloader_opts"],
stage=sb.Stage.TEST
)

# from speechbrain.dataio.feature_io import NumpyHdf5Reader
# reader = NumpyHdf5Reader(os.path.join(hparams["ssl_features_folder"], "train_hdf5_v2_ssl_feats.h5"))
# for item in train_data:
# print(item['id'])
# print(reader.read(item['id']))
# exit()
Original file line number Diff line number Diff line change
@@ -0,0 +1,81 @@
# ################################
# Model: wav2vec2
# Augmentation: SpecAugment
# Authors: Adel Moumen 2025
# ################################

# Seed needs to be set at top of yaml, before objects with parameters are made
seed: 1986
__set_seed: !apply:speechbrain.utils.seed_everything [!ref <seed>]
output_folder: !ref results/extract_ssl_representations/
save_folder: !ref <output_folder>/save
ssl_features_folder: !ref <output_folder>/ssl_features

# URL for the biggest Fairseq english wav2vec2 model.
wav2vec2_hub: /scratch/adelmou/models/facebook/wav2vec2-base-960h # /scratch/adelmou/models/facebook/hubert-large-ll60k #
wav2vec2_folder: !ref <save_folder>/wav2vec2_checkpoint

# Data files
data_folder: !PLACEHOLDER # e.g., /path/to/LibriSpeech
# noise/ris dataset will automatically be downloaded
# data_folder_rirs: !ref <data_folder>
train_splits: ["train-clean-100", "train-clean-360", "train-other-500"] # , "train-clean-360", "train-other-500"
dev_splits: ["dev-clean"]
test_splits: ["test-clean", "test-other"]
skip_prep: False
train_csv: !ref <output_folder>/train.csv
valid_csv: !ref <output_folder>/dev-clean.csv
test_csv:
- !ref <output_folder>/test-clean.csv
- !ref <output_folder>/test-other.csv

####################### Training Parameters ####################################

feature_configs:
ssl_feats: !new:speechbrain.dataio.feature_io.FeatureStorageConfig
name: ssl_feats
dtype: float32
writer_class: !name:speechbrain.dataio.feature_io.NumpyHdf5Writer

train_feature_storage_writers: !apply:speechbrain.dataio.feature_io.create_feature_storage_writers
feature_configs: !ref <feature_configs>
base_path: !ref <ssl_features_folder>
prefix: "train_960h"

valid_feature_storage_writers: !apply:speechbrain.dataio.feature_io.create_feature_storage_writers
feature_configs: !ref <feature_configs>
base_path: !ref <ssl_features_folder>
prefix: "dev_clean"

test_feature_storage_writers:
test-clean: !apply:speechbrain.dataio.feature_io.create_feature_storage_writers
feature_configs: !ref <feature_configs>
base_path: !ref <ssl_features_folder>
prefix: "test_clean"
test-other: !apply:speechbrain.dataio.feature_io.create_feature_storage_writers
feature_configs: !ref <feature_configs>
base_path: !ref <ssl_features_folder>
prefix: "test_other"

precision: bf16 # bf16, fp16 or fp32
sample_rate: 16000
freeze_wav2vec: True
blank_index: 0
# With data_parallel batch_size is split into N jobs
# With DDP batch_size is multiplied by N jobs
# Must be 3 per GPU to fit 32GB of VRAM
batch_size: 32

# Dataloader options
dataloader_opts:
batch_size: !ref <batch_size>


wav2vec2: !new:speechbrain.integrations.huggingface.wav2vec2.Wav2Vec2
source: !ref <wav2vec2_hub>
output_norm: True
freeze: !ref <freeze_wav2vec>
save_path: !ref <wav2vec2_folder>

modules:
wav2vec2: !ref <wav2vec2>
Loading