File size: 5,563 Bytes
6eed659
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
#!/usr/bin/env python3
"""Corrections POST-DECODAGE apprises sur le corpus WAXAL train (aucune donnee externe) :
  (A) RESPACING : le corpus lin joint parfois un mot-fonction au suivant ("namoni") la ou
      notre hypothese le separe ("na moni"), et inversement. On apprend, pour chaque paire
      (w1,w2), si le corpus prefere la forme jointe ou separee, et on applique la majorite.
  (B) NETTOYAGE des runs de lettres repetees (artefact type "aaana").
Evalue chaque correction isolement puis combinee, sur devhard-lin.
"""
import collections
import json
import pickle
import re

import jiwer
from multiprocessing import Pool
from pyctcdecode import build_ctcdecoder
from transformers import AutoProcessor

MDL = "/root/models/joint_cont_best"
ARPA = "/scratch/lm/lin_5g.arpa"
CORPUS = "/scratch/lm/corpus_lin.txt"


def comb(refs, hyps):
    pr = [(r, h) for r, h in zip(refs, hyps) if r.strip()]
    a = [x for x, _ in pr]
    b = [y for _, y in pr]
    w = jiwer.wer(a, b)
    c = jiwer.cer(a, b)
    return w, c, 0.5 * w + 0.5 * c


def build_stats():
    """joined[(w1,w2)] = nb de fois ou 'w1w2' apparait comme UN token dans le corpus.
    split[(w1,w2)]  = nb de fois ou 'w1 w2' apparait comme DEUX tokens consecutifs."""
    joined = collections.Counter()
    split = collections.Counter()
    unigram = collections.Counter()
    for line in open(CORPUS, encoding="utf-8"):
        toks = [t for t in re.split(r"\s+", line.strip()) if t]
        bare = [re.sub(r"^[^\w]+|[^\w]+$", "", t).lower() for t in toks]
        bare = [b for b in bare if b]
        unigram.update(bare)
        for a, b in zip(bare, bare[1:]):
            split[(a, b)] += 1
    # formes jointes : pour chaque token du corpus, tester les decoupages en 2 mots connus
    for w, n in unigram.items():
        for i in range(1, len(w)):
            a, b = w[:i], w[i:]
            if len(a) >= 2 and len(b) >= 2 and a in unigram and b in unigram:
                joined[(a, b)] += n
    return joined, split, unigram


def respace(text, joined, split, ratio):
    """Joint w1 w2 si le corpus prefere nettement la forme jointe."""
    toks = text.split()
    out = []
    i = 0
    while i < len(toks):
        if i + 1 < len(toks):
            a = re.sub(r"^[^\w]+|[^\w]+$", "", toks[i]).lower()
            b = re.sub(r"^[^\w]+|[^\w]+$", "", toks[i + 1]).lower()
            if a and b:
                j = joined.get((a, b), 0)
                s = split.get((a, b), 0)
                if j >= 3 and j > ratio * s:
                    merged = toks[i] + toks[i + 1]
                    out.append(merged)
                    i += 2
                    continue
        out.append(toks[i])
        i += 1
    return " ".join(out)


def declean(text):
    """Reduit les runs de la meme lettre a 2 max (artefact 'aaana' -> 'aana')."""
    return re.sub(r"(.)\1{2,}", r"\1\1", text)


def main():
    rows = [json.loads(l) for l in open("/root/devhard/devhard_linsna.jsonl", encoding="utf-8")]
    sub = [r for r in rows if r["lang"] == "lin"]
    refs = [r["text"] for r in sub]
    log = pickle.load(open("/scratch/lm/logits_lin.pkl", "rb"))
    tok = AutoProcessor.from_pretrained(MDL).tokenizer
    v = tok.get_vocab()
    lab = [None] * len(v)
    for t, i in v.items():
        lab[i] = t
    lab[tok.word_delimiter_token_id] = " "
    lab[tok.unk_token_id] = "⁇"
    lab[tok.pad_token_id] = ""
    greedy = [" ".join(tok.decode(l.argmax(-1)).replace("|", " ").split()) for l in log]
    dec = build_ctcdecoder(lab, kenlm_model_path=ARPA, alpha=0.5, beta=0.5,
                           lm_score_boundary=False)
    with Pool(8) as p:
        base = [" ".join(x.split()) for x in dec.decode_batch(p, log, beam_width=64)]
    base = [(g[:1] + h[1:] if h and g else h) for h, g in zip(base, greedy)]
    w0, c0, m0 = comb(refs, base)
    print("BASE (config optimale): combine=%.4f (WER %.4f CER %.4f)" % (m0, w0, c0), flush=True)

    joined, split, uni = build_stats()
    print("stats corpus: %d paires jointes, %d paires separees, %d unigrammes"
          % (len(joined), len(split), len(uni)), flush=True)

    print("\n--- (B) nettoyage des runs de lettres ---", flush=True)
    h = [declean(x) for x in base]
    _, _, m = comb(refs, h)
    print("  declean: %.4f  (%+.4f)" % (m, m - m0), flush=True)
    nchanged = sum(1 for a, b in zip(base, h) if a != b)
    print("  clips modifies: %d/%d" % (nchanged, len(base)), flush=True)

    print("\n--- (A) respacing appris, par seuil de ratio ---", flush=True)
    best = (m0, "base", None)
    for ratio in (0.5, 1.0, 2.0, 5.0):
        h = [respace(x, joined, split, ratio) for x in base]
        _, _, m = comb(refs, h)
        nch = sum(1 for a, b in zip(base, h) if a != b)
        print("  ratio=%.1f : %.4f  (%+.4f)  clips modifies=%d" % (ratio, m, m - m0, nch), flush=True)
        if m < best[0]:
            best = (m, "respace", ratio)

    print("\n--- combinaison des deux ---", flush=True)
    for ratio in (1.0, 2.0, 5.0):
        h = [declean(respace(x, joined, split, ratio)) for x in base]
        _, _, m = comb(refs, h)
        print("  respace(%.1f)+declean : %.4f  (%+.4f)" % (ratio, m, m - m0), flush=True)
        if m < best[0]:
            best = (m, "respace+declean", ratio)

    print("\nBEST_POST %.4f  %s ratio=%s  (base %.4f)" % (best[0], best[1], best[2], m0), flush=True)
    json.dump({"combine": best[0], "kind": best[1], "ratio": best[2], "base": m0},
              open("/root/respacing_best.json", "w"))
    print("RESPACING_DONE", flush=True)


if __name__ == "__main__":
    main()