sakha-punctuation / punctuate.py
loalkota's picture
Upload folder using huggingface_hub
374f4d8 verified
Raw
History Blame Contribute Delete
4.05 kB
"""Apply the punctuation/capitalisation tagger to a stream of ASR words.
Words arrive one per line, lowercase and unpunctuated, exactly as the acoustic
model emits them. They are processed in overlapping windows so that a word near
a window edge is still decided with context on both sides; the middle of each
window wins.
"""
import argparse, json, sys, torch
PUNCT_LABELS = ["NONE", "COMMA", "PERIOD", "QUESTION", "EXCLAM", "COLON",
"SEMICOLON", "DASH"]
CASE_LABELS = ["LOWER", "TITLE", "UPPER"]
PUNCT_TEXT = {"NONE": "", "COMMA": ",", "PERIOD": ".", "QUESTION": "?",
"EXCLAM": "!", "COLON": ":", "SEMICOLON": ";", "DASH": " —"}
# The tagger was trained on the model's alphabet; Sakha spelling comes back after
TO_BASE = str.maketrans({"ҕ": "ғ", "ҥ": "ң"})
class Tagger(torch.nn.Module):
def __init__(self, encoder, hidden):
super().__init__()
self.encoder = encoder
self.dropout = torch.nn.Dropout(0.1)
self.punct = torch.nn.Linear(hidden, len(PUNCT_LABELS))
self.case = torch.nn.Linear(hidden, len(CASE_LABELS))
def forward(self, input_ids, attention_mask):
h = self.encoder(input_ids=input_ids,
attention_mask=attention_mask).last_hidden_state
h = self.dropout(h)
return self.punct(h), self.case(h)
def main():
p = argparse.ArgumentParser()
p.add_argument("--model_dir", required=True)
p.add_argument("--words", required=True)
p.add_argument("--out", required=True)
p.add_argument("--window", type=int, default=120)
p.add_argument("--stride", type=int, default=80)
p.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
args = p.parse_args()
from transformers import AutoModel, AutoTokenizer
ck = torch.load(args.model_dir + "/punct_model.pt", map_location="cpu",
weights_only=False)
tok = AutoTokenizer.from_pretrained(ck["base"])
enc = AutoModel.from_pretrained(ck["base"])
model = Tagger(enc, enc.config.hidden_size)
model.load_state_dict(ck["state_dict"])
model.eval().to(args.device)
words = [w.strip() for w in open(args.words, encoding="utf-8") if w.strip()]
print("words: %d" % len(words), flush=True)
votes = [dict() for _ in words]
for start in range(0, len(words), args.stride):
chunk = words[start:start + args.window]
if not chunk:
break
e = tok([w.translate(TO_BASE) for w in chunk], is_split_into_words=True,
truncation=True, max_length=512, return_tensors="pt")
wid = e.word_ids()
with torch.inference_mode():
pl, cl = model(e["input_ids"].to(args.device),
e["attention_mask"].to(args.device))
pp, cc = pl.argmax(-1)[0].cpu(), cl.argmax(-1)[0].cpu()
prev = None
for t, w in enumerate(wid):
if w is None or w == prev:
prev = w
continue
prev = w
i = start + w
if i >= len(words):
continue
# distance from the window centre decides whose vote counts
d = abs(w - len(chunk) / 2)
if "d" not in votes[i] or d < votes[i]["d"]:
votes[i] = {"d": d, "punct": PUNCT_LABELS[pp[t]],
"case": CASE_LABELS[cc[t]]}
if start + args.window >= len(words):
break
out = []
for w, v in zip(words, votes):
out.append({"word": w, "punct": v.get("punct", "NONE"),
"case": v.get("case", "LOWER")})
json.dump(out, open(args.out, "w"), ensure_ascii=False)
def render(o):
s = o["word"]
if o["case"] == "TITLE":
s = s[:1].upper() + s[1:]
elif o["case"] == "UPPER":
s = s.upper()
return s + PUNCT_TEXT[o["punct"]]
print(" ".join(render(o) for o in out[:60]), flush=True)
print("\nwrote %s" % args.out)
if __name__ == "__main__":
sys.exit(main())