diff --git a/requirements.txt b/requirements.txt index b3e884f..3fb7ec7 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,13 +1,18 @@ torch==2.2.0 torchvision==0.17.0 torchaudio==2.2.0 -soundfile +sentencepiece==0.2.1 numpy==1.26.4 + +datasets +editdistance +soundfile torch_audiomentations tqdm matplotlib pandas wandb +wget comet_ml hydra-core diff --git a/src/configs/baseline.yaml b/src/configs/baseline.yaml index a671c0c..5026f35 100644 --- a/src/configs/baseline.yaml +++ b/src/configs/baseline.yaml @@ -1,33 +1,33 @@ defaults: - - model: baseline - - writer: wandb + - model: conformer + - writer: cometml - metrics: example - datasets: example - dataloader: example - transforms: example_only_instance + - text_encoder: ctc_text_encoder - _self_ optimizer: _target_: torch.optim.AdamW lr: 3e-4 lr_scheduler: _target_: torch.optim.lr_scheduler.OneCycleLR - max_lr: 1e-2 + max_lr: 1e-4 pct_start: 0.1 - steps_per_epoch: ${trainer.epoch_len} epochs: ${trainer.n_epochs} + steps_per_epoch: null anneal_strategy: cos loss_function: _target_: src.loss.CTCLossWrapper -text_encoder: - _target_: src.text_encoder.CTCTextEncoder + zero_infinity: True trainer: log_step: 50 - n_epochs: 50 - epoch_len: 200 + n_epochs: 10 + epoch_len: null device_tensors: ["spectrogram", "text_encoded"] # which tensors should be on device (ex. GPU) resume_from: null # null or path to the checkpoint dir with *.pth and config.yaml device: auto # device name or "auto" - override: False # if True, will override the previous run with the same name + override: True # if True, will override the previous run with the same name monitor: "min val_WER_(Argmax)" # "off" or "max/min metric_name", i.e. our goal is to maximize/minimize metric save_period: 5 # checkpoint each save_period epochs in addition to the best epoch early_stop: ${trainer.n_epochs} # epochs for early stopping diff --git a/src/configs/dataloader/example.yaml b/src/configs/dataloader/example.yaml index 71ea8d6..4bbd643 100644 --- a/src/configs/dataloader/example.yaml +++ b/src/configs/dataloader/example.yaml @@ -1,4 +1,4 @@ _target_: torch.utils.data.DataLoader -batch_size: 10 +batch_size: 16 num_workers: 2 pin_memory: True diff --git a/src/configs/datasets/onebatchtest.yaml b/src/configs/datasets/onebatchtest.yaml index 211ca48..c0785aa 100644 --- a/src/configs/datasets/onebatchtest.yaml +++ b/src/configs/datasets/onebatchtest.yaml @@ -5,8 +5,6 @@ train: max_text_length: 200 limit: 2 instance_transforms: ${transforms.instance_transforms.train} -# we filter partitions in one batch test to check the pipeline -# do not filter test dataset, you want to evaluate on the whole dataset val: _target_: src.datasets.LibrispeechDataset part: "dev-clean" diff --git a/src/configs/inference.yaml b/src/configs/inference.yaml index e252952..bf3e923 100644 --- a/src/configs/inference.yaml +++ b/src/configs/inference.yaml @@ -1,13 +1,14 @@ defaults: - - model: baseline - - metrics: example - - datasets: example_eval # we do not want to run inference on training data + - model: conformer + - metrics: beam_lm + - datasets: eval - dataloader: example - - transforms: example + - transforms: example_only_instance + - text_encoder: ctc_text_encoder - _self_ inferencer: - device_tensors: ["data_object", "labels"] # which tensors should be on device (ex. GPU) + device_tensors: ["spectrogram", "text_encoded"] # which tensors should be on device (ex. GPU) device: auto # device name or "auto" save_path: "example" # any name here, can be a dataset name seed: 1 - from_pretrained: "saved/model_best.pth" # path to the pretrained model + from_pretrained: "saved/beam_search_big_model/model_best.pth" # path to the pretrained model diff --git a/src/datasets/collate.py b/src/datasets/collate.py index 3f8a624..33b0bbc 100644 --- a/src/datasets/collate.py +++ b/src/datasets/collate.py @@ -1,17 +1,40 @@ import torch +from torch.nn.utils.rnn import pad_sequence def collate_fn(dataset_items: list[dict]): - """ - Collate and pad fields in the dataset items. - Converts individual items into a batch. - - Args: - dataset_items (list[dict]): list of objects from - dataset.__getitem__. - Returns: - result_batch (dict[Tensor]): dict, containing batch-version - of the tensors. - """ - - pass # TODO + audios = [it["audio"].squeeze(0) for it in dataset_items] + specs = [it["spectrogram"] for it in dataset_items] + texts = [it["text"] for it in dataset_items] + paths = [it["audio_path"] for it in dataset_items] + + txt_ids = [it["text_encoded"].squeeze(0).to(dtype=torch.long) + for it in dataset_items] + + audio_len = torch.tensor([a.numel() for a in audios], dtype=torch.long) + spec_len = torch.tensor([s.shape[-1] for s in specs], dtype=torch.long) + txt_len = torch.tensor([t.numel() for t in txt_ids], dtype=torch.long) + + audio_batch = pad_sequence(audios, batch_first=True, padding_value=0.0) + + specs_TF = [s.squeeze(0).transpose(0, 1) for s in specs] + specs_padded = pad_sequence(specs_TF, batch_first=True, padding_value=0.0) + spec_batch = specs_padded.transpose(1, 2).contiguous() + + text_encoded_batch = pad_sequence(txt_ids, batch_first=True, padding_value=0) + + batch = { + "audio": audio_batch.float(), + "audio_length": audio_len, + + "spectrogram": spec_batch.float(), + "spectrogram_length": spec_len, + + "text": texts, + "text_encoded": text_encoded_batch, + "text_encoded_length": txt_len, + + "audio_path": paths, + } + + return batch diff --git a/src/datasets/common_voice.py b/src/datasets/common_voice.py index 4ab0309..b94835b 100644 --- a/src/datasets/common_voice.py +++ b/src/datasets/common_voice.py @@ -12,7 +12,9 @@ from src.utils.io_utils import ROOT_PATH class CommonVoiceDataset(BaseDataset): def __init__(self, split, *args, **kwargs): - self._data_dir = ROOT_PATH / "dataset_common_voice" + self._data_dir = ROOT_PATH / "data"/ "datasets"/ "common_voice" + self._data_dir.mkdir(exist_ok=True, parents=True) + self._regex = re.compile("[^a-z ]") self._dataset = load_dataset( "mozilla-foundation/common_voice_11_0", diff --git a/src/logger/cometml.py b/src/logger/cometml.py index c97ad57..b3a10d6 100644 --- a/src/logger/cometml.py +++ b/src/logger/cometml.py @@ -62,6 +62,7 @@ class CometMLWriter: exp_class = comet_ml.Experiment self.exp = exp_class( + "XlK5HKtwRja1CCxz4bK6lfjCQ", project_name=project_name, workspace=workspace, experiment_key=self.run_id, diff --git a/src/logger/utils.py b/src/logger/utils.py index cdb60be..6855312 100644 --- a/src/logger/utils.py +++ b/src/logger/utils.py @@ -60,7 +60,7 @@ def plot_spectrogram(spectrogram, name=None): buf.seek(0) # convert buffer to Tensor - image = ToTensor()(PIL.Image.open(buf)) + image = ToTensor()(PIL.Image.open(buf).convert("RGB")) plt.close() diff --git a/src/metrics/__init__.py b/src/metrics/__init__.py index a8d047c..04c7eda 100644 --- a/src/metrics/__init__.py +++ b/src/metrics/__init__.py @@ -1,2 +1,2 @@ -from src.metrics.cer import ArgmaxCERMetric -from src.metrics.wer import ArgmaxWERMetric +from src.metrics.cer import ArgmaxCERMetric, BeamSearchCERMetric, CustomBeamSearchCERMetric +from src.metrics.wer import ArgmaxWERMetric, BeamSearchWERMetric, CustomBeamSearchWERMetric diff --git a/src/metrics/base_metric.py b/src/metrics/base_metric.py index f83ebcf..e9982cd 100644 --- a/src/metrics/base_metric.py +++ b/src/metrics/base_metric.py @@ -1,4 +1,12 @@ from abc import abstractmethod +import torch +from torch import Tensor +from typing import List + +from src.utils.io_utils import ROOT_PATH +from src.metrics.utils import expand_and_merge_beams, truncate_beams + +from pyctcdecode import build_ctcdecoder class BaseMetric: @@ -20,3 +28,99 @@ class BaseMetric: Can use external functions (like TorchMetrics) or custom ones. """ raise NotImplementedError() + + +class ArgmaxMetric(BaseMetric): + def __init__(self, text_encoder, binary_func=None, *args, **kwargs): + super().__init__(*args, **kwargs) + self.text_encoder = text_encoder + self.binary_func = binary_func + + def __call__( + self, log_probs: Tensor, log_probs_length: Tensor, text: List[str], **kwargs + ): + results = [] + predictions = torch.argmax(log_probs.detach().cpu(), dim=-1).numpy() + lengths = log_probs_length.detach().numpy() + for log_prob_vec, length, target_text in zip(predictions, lengths, text): + target_text = self.text_encoder.normalize_text(target_text) + pred_text = self.text_encoder.ctc_decode(log_prob_vec[:length]) + results.append(self.binary_func(target_text, pred_text)) + return sum(results) / len(results) + + +class BeamSearchMetric(BaseMetric): + def __init__(self, text_encoder, beam_size=20, binary_func=None, *args, **kwargs): + super().__init__(*args, **kwargs) + self.text_encoder = text_encoder + self.beam_size = beam_size + self.EMPTY_TOK = text_encoder.EMPTY_TOK + self.binary_func = binary_func + + ind2char = {i: self.text_encoder[i] for i in range(len(self.text_encoder))} + VOCAB = [ind2char[i] for i in range(len(ind2char))] + + ctcdecoder_args = {} + if "use_lm" in kwargs and kwargs["use_lm"]: + from torchaudio.models.decoder import download_pretrained_files + files = download_pretrained_files("librispeech-4-gram") + ctcdecoder_args["kenlm_model_path"] = files.lm + + if "alpha" in kwargs: + ctcdecoder_args["alpha"] = kwargs["alpha"] + if "beta" in kwargs: + ctcdecoder_args["beta"] = kwargs["beta"] + + self.decoder = build_ctcdecoder(VOCAB, **ctcdecoder_args) + + def __call__( + self, log_probs: Tensor, log_probs_length: Tensor, text: List[str], **kwargs, + ): + results = [] + for i, (T, target_text) in enumerate(zip(log_probs_length.detach().cpu().tolist(), text)): + logits = log_probs[i, :T, :].detach().cpu().numpy() + pred_text = self.decoder.decode(logits, beam_width=self.beam_size) + + if getattr(self.text_encoder, "SPACE_PIECE", False): + pred_text = pred_text.replace(self.text_encoder.SPACE_PIECE, " ") + + ref = self.text_encoder.normalize_text(target_text) + results.append(self.binary_func(ref, pred_text)) + return sum(results) / len(results) + + +class CustomBeamSearchMetric(BaseMetric): + def __init__(self, text_encoder, beam_size=20, binary_func=None, *args, **kwargs): + super().__init__(*args, **kwargs) + self.text_encoder = text_encoder + self.beam_size = beam_size + self.EMPTY_TOK = text_encoder.EMPTY_TOK + self.binary_func = binary_func + + ind2char = {i: self.text_encoder[i] for i in range(len(self.text_encoder))} + self.VOCAB = [ind2char[i] for i in range(len(ind2char))] + + def __call__( + self, log_probs: Tensor, log_probs_length: Tensor, text: List[str], **kwargs, + ): + results = [] + + for i, (T, target_text) in enumerate(zip(log_probs_length.detach().cpu().tolist(), text)): + probs = log_probs[i, :T, :].detach().cpu().exp() + dp = {("", self.EMPTY_TOK): 1.0} + for t in range(T): + cur_step_prob = probs[t] + dp = expand_and_merge_beams(dp, cur_step_prob, self.VOCAB, self.EMPTY_TOK) + dp = truncate_beams(dp, self.beam_size) + + hypos = [(pref, proba) for (pref, _), proba in dp.items()] + hypos.sort(key=lambda x: -x[1]) + pred_text = hypos[0][0] if hypos else "" + + if getattr(self.text_encoder, "SPACE_PIECE", False): + pred_text = pred_text.replace(self.text_encoder.SPACE_PIECE, " ") + + ref = self.text_encoder.normalize_text(target_text) + results.append(self.binary_func(ref, pred_text)) + + return sum(results) / len(results) diff --git a/src/metrics/cer.py b/src/metrics/cer.py index 77be3ea..05a2a0a 100644 --- a/src/metrics/cer.py +++ b/src/metrics/cer.py @@ -1,29 +1,17 @@ -from typing import List +from src.metrics.base_metric import ArgmaxMetric, BeamSearchMetric, CustomBeamSearchMetric +from src.metrics.utils import calc_cer -import torch -from torch import Tensor -from src.metrics.base_metric import BaseMetric -from src.metrics.utils import calc_cer +class ArgmaxCERMetric(ArgmaxMetric): + def __init__(self, *args, **kwargs): + super().__init__(*args, binary_func=calc_cer, **kwargs) -# TODO add beam search/lm versions -# Note: they can be written in a pretty way -# Note 2: overall metric design can be significantly improved +class BeamSearchCERMetric(BeamSearchMetric): + def __init__(self, *args, **kwargs): + super().__init__(*args, binary_func=calc_cer, **kwargs) -class ArgmaxCERMetric(BaseMetric): - def __init__(self, text_encoder, *args, **kwargs): - super().__init__(*args, **kwargs) - self.text_encoder = text_encoder - def __call__( - self, log_probs: Tensor, log_probs_length: Tensor, text: List[str], **kwargs - ): - cers = [] - predictions = torch.argmax(log_probs.cpu(), dim=-1).numpy() - lengths = log_probs_length.detach().numpy() - for log_prob_vec, length, target_text in zip(predictions, lengths, text): - target_text = self.text_encoder.normalize_text(target_text) - pred_text = self.text_encoder.ctc_decode(log_prob_vec[:length]) - cers.append(calc_cer(target_text, pred_text)) - return sum(cers) / len(cers) +class CustomBeamSearchCERMetric(CustomBeamSearchMetric): + def __init__(self, *args, **kwargs): + super().__init__(*args, binary_func=calc_cer, **kwargs) diff --git a/src/metrics/utils.py b/src/metrics/utils.py index 88d48b9..3307675 100644 --- a/src/metrics/utils.py +++ b/src/metrics/utils.py @@ -1,13 +1,44 @@ -# Based on seminar materials +import editdistance +from collections import defaultdict +import tempfile +import subprocess -# Don't forget to support cases when target_text == '' +from src.utils.io_utils import ROOT_PATH + + +def calc_wer(target_text, predicted_text) -> float: + if not target_text: + return 1 if predicted_text else 0 + return editdistance.eval(target_text.split(), predicted_text.split()) / len(target_text.split()) def calc_cer(target_text, predicted_text) -> float: - # TODO - pass + if not target_text: + return 1 if predicted_text else 0 + return editdistance.eval(target_text, predicted_text) / len(target_text) -def calc_wer(target_text, predicted_text) -> float: - # TODO - pass +def expand_and_merge_beams(dp, cur_step_prob, vocab, empty_tok): + new_dp = defaultdict(float) + + for (pref, prev_char), pref_proba in dp.items(): + for idx, char in enumerate(vocab): + cur_proba = pref_proba * cur_step_prob[idx] + cur_char = char + + if char == empty_tok: + cur_pref = pref + else: + if prev_char != char: + cur_pref = pref + char + else: + cur_pref = pref + + new_dp[(cur_pref, cur_char)] += cur_proba + + return new_dp + + +def truncate_beams(dp, beam_size): + items = sorted(list(dp.items()), key=lambda x: -x[1])[:beam_size] + return dict(items) diff --git a/src/metrics/wer.py b/src/metrics/wer.py index da80390..251cbc3 100644 --- a/src/metrics/wer.py +++ b/src/metrics/wer.py @@ -1,29 +1,17 @@ -from typing import List +from src.metrics.base_metric import ArgmaxMetric, BeamSearchMetric, CustomBeamSearchMetric +from src.metrics.utils import calc_wer -import torch -from torch import Tensor -from src.metrics.base_metric import BaseMetric -from src.metrics.utils import calc_wer +class ArgmaxWERMetric(ArgmaxMetric): + def __init__(self, *args, **kwargs): + super().__init__(*args, binary_func=calc_wer, **kwargs) -# TODO beam search / LM versions -# Note: they can be written in a pretty way -# Note 2: overall metric design can be significantly improved +class BeamSearchWERMetric(BeamSearchMetric): + def __init__(self, *args, **kwargs): + super().__init__(*args, binary_func=calc_wer, **kwargs) -class ArgmaxWERMetric(BaseMetric): - def __init__(self, text_encoder, *args, **kwargs): - super().__init__(*args, **kwargs) - self.text_encoder = text_encoder - def __call__( - self, log_probs: Tensor, log_probs_length: Tensor, text: List[str], **kwargs - ): - wers = [] - predictions = torch.argmax(log_probs.cpu(), dim=-1).numpy() - lengths = log_probs_length.detach().numpy() - for log_prob_vec, length, target_text in zip(predictions, lengths, text): - target_text = self.text_encoder.normalize_text(target_text) - pred_text = self.text_encoder.ctc_decode(log_prob_vec[:length]) - wers.append(calc_wer(target_text, pred_text)) - return sum(wers) / len(wers) +class CustomBeamSearchWERMetric(CustomBeamSearchMetric): + def __init__(self, *args, **kwargs): + super().__init__(*args, binary_func=calc_wer, **kwargs) diff --git a/src/model/__init__.py b/src/model/__init__.py index 014317e..9425740 100644 --- a/src/model/__init__.py +++ b/src/model/__init__.py @@ -1,5 +1,7 @@ from src.model.baseline_model import BaselineModel +from src.model.conformer_model import ConformerModel __all__ = [ "BaselineModel", + "ConformerModel", ] diff --git a/src/text_encoder/__init__.py b/src/text_encoder/__init__.py index b66bd53..d34f106 100644 --- a/src/text_encoder/__init__.py +++ b/src/text_encoder/__init__.py @@ -1 +1,2 @@ from src.text_encoder.ctc_text_encoder import CTCTextEncoder +from src.text_encoder.sp_ctc_text_encoder import SPCTCTextEncoder diff --git a/src/text_encoder/ctc_text_encoder.py b/src/text_encoder/ctc_text_encoder.py index 1d66445..17e239e 100644 --- a/src/text_encoder/ctc_text_encoder.py +++ b/src/text_encoder/ctc_text_encoder.py @@ -3,12 +3,6 @@ from string import ascii_lowercase import torch -# TODO add CTC decode -# TODO add BPE, LM, Beam Search support -# Note: think about metrics and encoder -# The design can be remarkably improved -# to calculate stuff more efficiently and prettier - class CTCTextEncoder: EMPTY_TOK = "" @@ -59,7 +53,23 @@ class CTCTextEncoder: return "".join([self.ind2char[int(ind)] for ind in inds]).strip() def ctc_decode(self, inds) -> str: - pass # TODO + seq = inds.tolist() + + blank_id = self.char2ind[self.EMPTY_TOK] + result_chars = "" + prev = None + + for idx in seq: + if idx == blank_id: + prev = idx + continue + if idx == prev: + continue + ch = self.ind2char[idx] + result_chars += ch + prev = idx + + return result_chars.strip() @staticmethod def normalize_text(text: str): diff --git a/src/trainer/inferencer.py b/src/trainer/inferencer.py index 381a31f..26e6e64 100644 --- a/src/trainer/inferencer.py +++ b/src/trainer/inferencer.py @@ -136,26 +136,24 @@ class Inferencer(BaseTrainer): # Some saving logic. This is an example # Use if you need to save predictions on disk - batch_size = batch["logits"].shape[0] - current_id = batch_idx * batch_size - - for i in range(batch_size): - # clone because of - # https://github.com/pytorch/pytorch/issues/1995 - logits = batch["logits"][i].clone() - label = batch["labels"][i].clone() - pred_label = logits.argmax(dim=-1) - - output_id = current_id + i - - output = { - "pred_label": pred_label, - "label": label, - } - - if self.save_path is not None: - # you can use safetensors or other lib here - torch.save(output, self.save_path / part / f"output_{output_id}.pth") + # batch_size = batch["log_probs"].shape[0] + # current_id = batch_idx * batch_size + + # for i in range(batch_size): + # logits = batch["log_probs"][i].clone() + # label = batch["labels"][i].clone() + # pred_label = logits.argmax(dim=-1) + + # output_id = current_id + i + + # output = { + # "pred_label": pred_label, + # "label": label, + # } + + # if self.save_path is not None: + # # you can use safetensors or other lib here + # torch.save(output, self.save_path / part / f"output_{output_id}.pth") return batch diff --git a/src/trainer/trainer.py b/src/trainer/trainer.py index d0ef8fa..1fbdb2c 100644 --- a/src/trainer/trainer.py +++ b/src/trainer/trainer.py @@ -86,7 +86,7 @@ class Trainer(BaseTrainer): def log_spectrogram(self, spectrogram, **batch): spectrogram_for_plot = spectrogram[0].detach().cpu() - image = plot_spectrogram(spectrogram_for_plot) + image = plot_spectrogram(spectrogram_for_plot).permute(1, 2, 0) self.writer.add_image("spectrogram", image) def log_predictions( diff --git a/train.py b/train.py index e7e6c08..628ba17 100644 --- a/train.py +++ b/train.py @@ -57,12 +57,20 @@ def main(config): # build optimizer, learning rate scheduler trainable_params = filter(lambda p: p.requires_grad, model.parameters()) + optimizer = instantiate(config.optimizer, params=trainable_params) - lr_scheduler = instantiate(config.lr_scheduler, optimizer=optimizer) # epoch_len = number of iterations for iteration-based training # epoch_len = None or len(dataloader) for epoch-based training - epoch_len = config.trainer.get("epoch_len") + epoch_len = config.trainer.get("epoch_len") or len(dataloaders["train"]) + + sched_kwargs = { + "optimizer": optimizer, + } + if "steps_per_epoch" in config.lr_scheduler: + sched_kwargs["steps_per_epoch"] = epoch_len + + lr_scheduler = instantiate(config.lr_scheduler, **sched_kwargs) trainer = Trainer( model=model,