asr_checkpoint_baseline / git_diff.patch
ZenMan67's picture
Upload folder using huggingface_hub
36b66b3 verified
Raw
History Blame Contribute Delete
17.7 kB
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/inference.yaml b/src/configs/inference.yaml
index e252952..0499056 100644
--- a/src/configs/inference.yaml
+++ b/src/configs/inference.yaml
@@ -4,6 +4,7 @@ defaults:
- datasets: example_eval # we do not want to run inference on training data
- dataloader: example
- transforms: example
+ - text_encoder: ctc_text_encoder
- _self_
inferencer:
device_tensors: ["data_object", "labels"] # which tensors should be on device (ex. GPU)
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..70a83fe 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
+from src.metrics.wer import ArgmaxWERMetric, BeamSearchWERMetric
diff --git a/src/metrics/cer.py b/src/metrics/cer.py
index 77be3ea..0af960d 100644
--- a/src/metrics/cer.py
+++ b/src/metrics/cer.py
@@ -4,9 +4,11 @@ import torch
from torch import Tensor
from src.metrics.base_metric import BaseMetric
-from src.metrics.utils import calc_cer
+from src.metrics.utils import calc_cer, _expand_and_merge_beams, _truncate_beams
-# TODO add beam search/lm versions
+from pyctcdecode import build_ctcdecoder
+
+# TODO lm versions
# Note: they can be written in a pretty way
# Note 2: overall metric design can be significantly improved
@@ -20,10 +22,51 @@ class ArgmaxCERMetric(BaseMetric):
self, log_probs: Tensor, log_probs_length: Tensor, text: List[str], **kwargs
):
cers = []
- predictions = torch.argmax(log_probs.cpu(), dim=-1).numpy()
+ 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])
cers.append(calc_cer(target_text, pred_text))
return sum(cers) / len(cers)
+
+
+class BeamSearchCERMetric(BaseMetric):
+ def __init__(self, text_encoder, beam_size=20, *args, **kwargs):
+ super().__init__(*args, **kwargs)
+ self.text_encoder = text_encoder
+ self.beam_size = beam_size
+ self.EMPTY_TOK = getattr(text_encoder, "EMPTY_TOK", "")
+
+ ind2char = {i: self.text_encoder[i] for i in range(len(self.text_encoder))}
+ VOCAB = [ind2char[i] for i in range(len(ind2char))]
+ self.decoder = build_ctcdecoder(VOCAB)
+
+ def __call__(
+ self, log_probs: Tensor, log_probs_length: Tensor, text: List[str], **kwargs,
+ ):
+ cers = []
+
+ ind2char = {i: self.text_encoder[i] for i in range(len(self.text_encoder))}
+ VOCAB = [ind2char[i] for i in range(len(ind2char))]
+
+ 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, 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 ""
+
+ logits = log_probs[i, :T, :].detach().cpu().numpy()
+ pred_text = self.decoder.decode(logits, beam_width=self.beam_size)
+
+ ref = self.text_encoder.normalize_text(target_text)
+ cers.append(calc_cer(ref, pred_text))
+
+ return sum(cers) / len(cers)
\ No newline at end of file
diff --git a/src/metrics/utils.py b/src/metrics/utils.py
index 88d48b9..e13b32b 100644
--- a/src/metrics/utils.py
+++ b/src/metrics/utils.py
@@ -1,13 +1,40 @@
-# Based on seminar materials
+import editdistance
+from collections import defaultdict
-# Don't forget to support cases when target_text == ''
+
+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..c4e7de6 100644
--- a/src/metrics/wer.py
+++ b/src/metrics/wer.py
@@ -4,9 +4,11 @@ import torch
from torch import Tensor
from src.metrics.base_metric import BaseMetric
-from src.metrics.utils import calc_wer
+from src.metrics.utils import calc_wer, _expand_and_merge_beams, _truncate_beams
-# TODO beam search / LM versions
+from pyctcdecode import build_ctcdecoder
+
+# TODO LM versions
# Note: they can be written in a pretty way
# Note 2: overall metric design can be significantly improved
@@ -21,9 +23,58 @@ class ArgmaxWERMetric(BaseMetric):
):
wers = []
predictions = torch.argmax(log_probs.cpu(), dim=-1).numpy()
- lengths = log_probs_length.detach().numpy()
+ lengths = log_probs_length.detach().cpu().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 BeamSearchWERMetric(BaseMetric):
+ def __init__(self, text_encoder, beam_size=20, *args, **kwargs):
+ super().__init__(*args, **kwargs)
+ self.text_encoder = text_encoder
+ self.beam_size = beam_size
+ self.EMPTY_TOK = getattr(text_encoder, "EMPTY_TOK", "")
+
+ ind2char = {i: self.text_encoder[i] for i in range(len(self.text_encoder))}
+ VOCAB = [ind2char[i] for i in range(len(ind2char))]
+ self.decoder = build_ctcdecoder(VOCAB)
+
+ def __call__(
+ self, log_probs: Tensor, log_probs_length: Tensor, text: List[str], **kwargs,
+ ):
+ wers = []
+
+ ind2char = {i: self.text_encoder[i] for i in range(len(self.text_encoder))}
+ VOCAB = [ind2char[i] for i in range(len(ind2char))]
+
+ 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}
+ # print(i, T)
+ # import time
+ # expand_time, truncate_time = 0, 0
+ # for t in range(T):
+ # cur_step_prob = probs[t]
+ # t1 = time.perf_counter()
+ # dp = _expand_and_merge_beams(dp, cur_step_prob, VOCAB, self.EMPTY_TOK)
+ # t2 = time.perf_counter()
+ # expand_time += t2 - t1
+ # dp = _truncate_beams(dp, self.beam_size)
+ # t3 = time.perf_counter()
+ # truncate_time += t3 - t2
+ # print(expand_time, truncate_time)
+
+ # hypos = [(pref, proba) for (pref, _), proba in dp.items()]
+ # hypos.sort(key=lambda x: -x[1])
+ # pred_text = hypos[0][0] if hypos else ""
+
+ logits = log_probs[i, :T, :].detach().cpu().numpy()
+ pred_text = self.decoder.decode(logits, beam_width=self.beam_size)
+
+ ref = self.text_encoder.normalize_text(target_text)
+ wers.append(calc_wer(ref, pred_text))
+
+ return sum(wers) / len(wers)
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/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,