File size: 4,133 Bytes
ce3c8df
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
"""
WER evaluation on a LibriSpeech split, via greedy CTC and greedy attention decode.

Usage (from project root, venv active):
    python -m src.evaluate --config configs/zipformer_s.yaml \
        --checkpoint checkpoints/latest.pt --split test-clean
"""

import os

os.environ.setdefault("PYTORCH_ENABLE_MPS_FALLBACK", "1")

import argparse
import warnings

import jiwer
import torch
import yaml
from torch.utils.data import DataLoader
from tqdm import tqdm

warnings.filterwarnings("ignore", message=".*output with one or more elements was resized.*")

from src.dataset import ASRCollate, LibriSpeechASR
from src.model import ASRModel
from src.tokenizer import ASRTokenizer, BOS_ID, EOS_ID


def ctc_greedy_decode(log_probs: torch.Tensor, lengths: torch.Tensor, blank_id: int) -> list:
    """log_probs: (B, T, V). Returns a list of token-id lists, one per sample."""
    ids = log_probs.argmax(dim=-1)  # (B, T)
    results = []
    for i in range(ids.size(0)):
        length = int(lengths[i].item())
        seq = ids[i, :length].tolist()
        collapsed = []
        prev = None
        for tok in seq:
            if tok != prev and tok != blank_id:
                collapsed.append(tok)
            prev = tok
        results.append(collapsed)
    return results


@torch.no_grad()
def evaluate(model: ASRModel, tokenizer: ASRTokenizer, loader: DataLoader, device: torch.device, max_batches=None):
    model.eval()
    refs = []
    ctc_hyps = []
    attn_hyps = []

    for i, batch in enumerate(tqdm(loader, desc="evaluating")):
        if max_batches is not None and i >= max_batches:
            break
        waveforms = batch["waveforms"].to(device)
        wave_lengths = batch["wave_lengths"].to(device)

        enc_out, enc_lengths, ctc_log_probs = model.forward_eval(waveforms, wave_lengths)

        ctc_ids = ctc_greedy_decode(ctc_log_probs, enc_lengths, blank_id=model.blank_id)
        attn_ids = model.decoder.greedy_decode(enc_out, enc_lengths, bos_id=BOS_ID, eos_id=EOS_ID, max_len=200)

        for ref, ctc_id_seq, attn_id_seq in zip(batch["transcripts"], ctc_ids, attn_ids):
            refs.append(ref.lower())
            ctc_hyps.append(tokenizer.decode(ctc_id_seq))
            attn_hyps.append(tokenizer.decode(attn_id_seq))

    ctc_wer = jiwer.wer(refs, ctc_hyps)
    attn_wer = jiwer.wer(refs, attn_hyps)
    return {"ctc_wer": ctc_wer, "attn_wer": attn_wer, "refs": refs, "ctc_hyps": ctc_hyps, "attn_hyps": attn_hyps}


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--config", required=True)
    parser.add_argument("--checkpoint", required=True)
    parser.add_argument("--split", default="test-clean")
    parser.add_argument("--max-batches", type=int, default=None)
    parser.add_argument("--batch-size", type=int, default=None)
    parser.add_argument("--show-examples", type=int, default=5)
    args = parser.parse_args()

    with open(args.config) as f:
        cfg = yaml.safe_load(f)

    device = torch.device("mps" if torch.backends.mps.is_available() else "cpu")

    tokenizer = ASRTokenizer(cfg["tokenizer_model"])
    model = ASRModel(vocab_size=tokenizer.vocab_size, **cfg["model"]).to(device)

    ckpt = torch.load(args.checkpoint, map_location=device)
    model.load_state_dict(ckpt["model"])
    print(f"Loaded checkpoint {args.checkpoint} (epoch {ckpt.get('epoch')}, step {ckpt.get('step')})")

    ds = LibriSpeechASR(cfg["data_root"], [args.split], download=False)
    collate = ASRCollate(tokenizer)
    batch_size = args.batch_size or cfg["batch_size"]
    loader = DataLoader(ds, batch_size=batch_size, shuffle=False, collate_fn=collate)

    results = evaluate(model, tokenizer, loader, device, max_batches=args.max_batches)

    print(f"\nCTC (greedy) WER:       {results['ctc_wer'] * 100:.2f}%")
    print(f"Attention (greedy) WER: {results['attn_wer'] * 100:.2f}%")

    n = min(args.show_examples, len(results["refs"]))
    for i in range(n):
        print(f"\nREF : {results['refs'][i]}")
        print(f"CTC : {results['ctc_hyps'][i]}")
        print(f"ATTN: {results['attn_hyps'][i]}")


if __name__ == "__main__":
    main()