{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Fine Tuning of Speech Recognition Transformer (HF)","metadata":{"id":"g4JTHJ-IATho"}},{"cell_type":"code","source":"%%capture\n!pip install datasets==1.14\n!pip install transformers==4.11.3\n!pip install librosa\n!apt install git-lfs","metadata":{"id":"L1532RVbJgQV"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport random\nimport numpy as np\n\nimport librosa\nfrom scipy.io.wavfile import write\n\nimport datasets\nfrom datasets import load_dataset, load_metric, Dataset, concatenate_datasets, Audio\n\nfrom tqdm import tqdm\nfrom IPython.display import Audio, display","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Fine-tuning a model on an audio classification task","metadata":{"id":"XalxdrirGkLl"}},{"cell_type":"markdown","source":"### Load datasets from kaggle","metadata":{"id":"mcE455KaG687"}},{"cell_type":"code","source":"# I was training on Google Colab Pro so i mount google drive\nfrom google.colab import drive\ndrive.mount('/content/drive')","metadata":{"id":"t6697XZrgRI6","outputId":"db8b4a08-6cdb-4978-ce48-1770cead5982"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install opendatasets\nimport opendatasets as od\nod.download('https://www.kaggle.com/c/classification-of-short-noisy-audio-speech/data')\nod.download('https://www.kaggle.com/chrisfilo/urbansound8k')\n!rm urbansound8k/UrbanSound8K.csv","metadata":{"id":"Dg8JP5d4NkXc"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Also could train facebook/hubert-base\nmodel_checkpoint = \"facebook/wav2vec-base\"\nbatch_size = 128\nSAMPLING_RATE = 16000 # Typical for speech recognition task sampling rate\nlabels = ['down', 'go', 'left', 'no', 'off', 'on', 'right', 'stop', 'up', 'yes']","metadata":{"id":"5WMEawzyCEyG"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Dataset preparing","metadata":{}},{"cell_type":"code","source":"# Download some noise recordings per each class for noising clear audio\nMAX_NOISE_NUMBER_PER_CLASS = 100\n\nhackathon_noise_path = 'classification-of-short-noisy-audio-speech/hackaton_ds/noises'\nnoise_folders = [os.path.join('urbansound8k', folder_name) for folder_name in os.listdir('urbansound8k')] + [hackathon_noise_path]\nnoises = []\nfor noise_folder in noise_folders:\n    noises.extend([os.path.join(noise_folder, x) for x in os.listdir(noise_folder)[:MAX_NOISE_NUMBER_PER_CLASS]])\nrandom.shuffle(noises)","metadata":{"id":"7WZtc-evUXsm"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# This cell could running for more than 10 minutes, that's ok\nnoises = [librosa.load(noise, sr=SAMPLING_RATE)[0] for noise in noises]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# This function is given by TV Neuro Technologies competition\ndef make_some_noise(clean):\n    ok = False\n    # Not all noises are good for our task so if something wrong happened just try one more time :D\n    while not ok:\n        try:\n            clean = np.array(clean)\n            max_amp = 1 - np.random.random() / 3\n            noise = noises[np.random.randint(0, len(noises))]\n            noise_amp = np.random.rand() * max_amp\n            max_start = len(noise) - len(clean)\n            start = np.random.randint(0, max_start + 1)\n            noise_part = noise[start:start+len(clean)]\n            noise_mult = np.abs(clean.max()) / (np.abs(noise_part).max() * noise_amp + 1e-9)\n            ok = True\n        except:\n            ok = False\n    return (clean + noise_part * noise_mult) / (1 + noise_amp + 1e-9)","metadata":{"id":"4WXXtNy6UVdv"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"save_path = 'classification-of-short-noisy-audio-speech/hackaton_ds/generated'\n!rm -r 'classification-of-short-noisy-audio-speech/hackaton_ds/generated'\n!mkdir 'classification-of-short-noisy-audio-speech/hackaton_ds/generated'\n\n!mkdir 'classification-of-short-noisy-audio-speech/hackaton_ds/generated/yes'\n!mkdir 'classification-of-short-noisy-audio-speech/hackaton_ds/generated/no'\n!mkdir 'classification-of-short-noisy-audio-speech/hackaton_ds/generated/up'\n!mkdir 'classification-of-short-noisy-audio-speech/hackaton_ds/generated/down'\n!mkdir 'classification-of-short-noisy-audio-speech/hackaton_ds/generated/left'\n!mkdir 'classification-of-short-noisy-audio-speech/hackaton_ds/generated/right'\n!mkdir 'classification-of-short-noisy-audio-speech/hackaton_ds/generated/on'\n!mkdir 'classification-of-short-noisy-audio-speech/hackaton_ds/generated/off'\n!mkdir 'classification-of-short-noisy-audio-speech/hackaton_ds/generated/stop'\n!mkdir 'classification-of-short-noisy-audio-speech/hackaton_ds/generated/go'","metadata":{"id":"LL5ancqCVqh7"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# superb-ks dataset has too clear records, let's noise them\ndataset = load_dataset(\"superb\", \"ks\")\nfor dataset_type in ['train', 'validation', 'test']:\n    current_dataset = dataset[dataset_type]\n    for i, x in enumerate(tqdm(current_dataset)):\n        label = x['file'].split('/')[-2]\n        if label in labels:\n            audio, _ = librosa.load(x['audio'], sr=SAMPLING_RATE)\n            audio = make_some_noise(audio)\n            write(f'{save_path}/{label}/{dataset_type[1]}-{i}.wav', SAMPLING_RATE, audio)","metadata":{"id":"YMzrMAcNTXr2","outputId":"c00424f4-b36d-4f25-d752-761f2b2b95cc"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DATA_PATH = 'classification-of-short-noisy-audio-speech/hackaton_ds/train'\nDATA_PATH2 = 'classification-of-short-noisy-audio-speech/hackaton_ds/generated'\ndata = {\n    'train':{\n        'audio':[],\n        'file':[],\n        'label':[]\n    }, \n    'test':{\n        'audio':[],\n        'file':[],\n        'label':[]\n    }\n}\n\nfolder_paths = [os.path.join(DATA_PATH, folder_name) for folder_name in os.listdir(DATA_PATH)]\nfolder_paths += [os.path.join(DATA_PATH2, folder_name) for folder_name in os.listdir(DATA_PATH2) if folder_name in labels]\ntrain_paths = []\nval_paths = []\n\n# We take validation ONLY from TV Neuro Technologies data, because it's the main target\nval_size = 0.2\n\nfor folder_path in folder_paths:\n    filenames = os.listdir(folder_path)\n    paths = [os.path.join(folder_path, file_name) for file_name in filenames]\n    train_size = int(len(paths) * (1 - val_size))\n    _train_paths, _val_paths = paths[:train_size], paths[train_size:]\n    train_paths.extend(_train_paths)\n    if 'train' in folder_path:\n        val_paths.extend(_val_paths)\nrandom.shuffle(train_paths)\nrandom.shuffle(val_paths)\n\nfor train_path in tqdm(train_paths):\n    data['train']['audio'].append(train_path)\n    data['train']['file'].append(train_path)\n    data['train']['label'].append(train_path.split('/')[-2])\n\nfor test_path in tqdm(val_paths):\n    data['test']['audio'].append(test_path)\n    data['test']['file'].append(test_path)\n    data['test']['label'].append(test_path.split('/')[-2])","metadata":{"id":"CSj-Cur6NeFR","outputId":"cc7d40a2-c58b-4a66-dc7c-820f99e14d0f"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"features = datasets.Features(\n    {\n        \"audio\": datasets.Value(dtype='string', id=None),\n        \"file\": datasets.Value(dtype='string', id=None),\n        \"label\": datasets.ClassLabel(num_classes=10, names=labels, names_file=None, id=None),\n    }\n)\n\ntrain_dataset = Dataset.from_dict(data['train'], features=features)\ntest_dataset = Dataset.from_dict(data['test'], features=features)\n\ntrain_dataset = train_dataset.cast_column(\"audio\", Audio(sampling_rate=16000, mono=True, id=None))\ntest_dataset = test_dataset.cast_column(\"audio\", Audio(sampling_rate=16000, mono=True, id=None))\n\ndataset = load_dataset(\"superb\", \"ks\")\ndataset[\"train\"] = train_dataset\ndataset[\"test\"] = test_dataset\ndataset[\"validation\"] = test_dataset\nmetric = load_metric(\"accuracy\")","metadata":{"id":"bq6XtsPJdHgk","outputId":"6e380e9b-e381-4c87-cbf5-f61cb8ce8cf9"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label2id, id2label = dict(), dict()\nfor i, label in enumerate(labels):\n    label2id[label] = str(i)\n    id2label[str(i)] = label","metadata":{"id":"UuyXDtQqNUZW"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import random\nfrom IPython.display import Audio, display\n\nfor _ in range(5):\n    rand_idx = random.randint(0, len(dataset[\"train\"])-1)\n    example = dataset[\"train\"][rand_idx]\n    audio = example[\"audio\"]\n\n    print(f'Label: {id2label[str(example[\"label\"])]}')\n    print(f'Shape: {audio[\"array\"].shape}, sampling rate: {audio[\"sampling_rate\"]}')\n    display(Audio(audio[\"array\"], rate=audio[\"sampling_rate\"]))\n    print()","metadata":{"id":"RI0kb4woNvuU","outputId":"c73fdddd-8f13-4347-f0a7-509abcc0fb40"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Preprocessing the data","metadata":{"id":"4zxoikSOjs0K"}},{"cell_type":"code","source":"from transformers import AutoFeatureExtractor\n\nfeature_extractor = AutoFeatureExtractor.from_pretrained(model_checkpoint)","metadata":{"id":"G1bX4lGAO_d9"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def preprocess_function(examples):\n    audio_arrays = [x[\"array\"] for x in examples[\"audio\"]]\n    inputs = feature_extractor(\n        audio_arrays, \n        sampling_rate=feature_extractor.sampling_rate, \n        max_length=int(feature_extractor.sampling_rate * 1.0), \n        truncation=True, \n    )\n    return inputs","metadata":{"id":"4O_p3WrpRyej"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoded_dataset = dataset.map(preprocess_function, remove_columns=[\"audio\", \"file\"], batched=True)","metadata":{"id":"FfxFgJ_6qPPy"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training the model","metadata":{"id":"HOXmyPQ76Qv9"}},{"cell_type":"code","source":"from transformers import AutoModelForAudioClassification, TrainingArguments, Trainer\n\nnum_labels = len(id2label)\nmodel = AutoModelForAudioClassification.from_pretrained(\n    model_checkpoint, \n    num_labels=num_labels,\n    label2id=label2id,\n    id2label=id2label,\n)\n","metadata":{"id":"X9DDujL0q1ac"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"save_folder = f\"/content/drive/MyDrive/DataScience/Practice/TV NEUROTECH/TV_NEUROTECH_PRIVATE/checkpoints/wav2vec\"\n\nargs = TrainingArguments(\n    save_folder,\n    evaluation_strategy = \"epoch\",\n    save_strategy = \"epoch\",\n    learning_rate=1e-4,\n    per_device_train_batch_size=batch_size,\n    gradient_accumulation_steps=4,\n    per_device_eval_batch_size=batch_size,\n    num_train_epochs=15,\n    warmup_ratio=0.1,\n    logging_steps=10,\n    load_best_model_at_end=True,\n    metric_for_best_model=\"accuracy\",\n    push_to_hub=False,\n    fp16=True, # because we have weight limit and want best accuracy with fp16\n    dataloader_num_workers=4\n)","metadata":{"id":"xc_MTm0Ks3DF","outputId":"52b9528e-1701-4662-f071-fdb2e29b5e2d"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def compute_metrics(eval_pred):\n    predictions = np.argmax(eval_pred.predictions, axis=1)\n    return metric.compute(predictions=predictions, references=eval_pred.label_ids)","metadata":{"id":"EVWfiBuv2uCS"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = Trainer(\n    model,\n    args,\n    train_dataset=encoded_dataset[\"train\"],\n    eval_dataset=encoded_dataset[\"validation\"],\n    tokenizer=feature_extractor,\n    compute_metrics=compute_metrics,\n)","metadata":{"id":"McVoaCPr3Cj-","outputId":"90f6f25e-9cdb-4df9-8cf0-245f930ca878"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.train()","metadata":{"id":"Pps61vF_4QaH","outputId":"372446e7-e191-4e31-a0ca-4a36a3a2abf0"},"execution_count":null,"outputs":[]}]}