LGTM / lgtm /text.py
qnguyen3's picture
LGTM: PyTorch + ONNX weights and inference code
409d4fb verified
Raw History Blame Contribute Delete
3.54 kB
"""Text frontend: NFKD normalisation, light cleanup, "<lang>text</lang>" tagging, and a
code-point -> id lookup (unicode_indexer.json, 8322 ids)."""
import json
import re
from unicodedata import normalize
AVAILABLE_LANGS = ["en", "ko", "ja", "ar", "bg", "cs", "da", "de", "el", "es", "et", "fi", "fr", "hi", "hr", "hu", "id", "it", "lt", "lv", "nl", "pl", "pt", "ro", "ru", "sk", "sl", "sv", "tr", "uk", "vi", "na"]
_EMOJI = re.compile(
"[\U0001f600-\U0001f64f\U0001f300-\U0001f5ff\U0001f680-\U0001f6ff\U0001f700-\U0001f77f"
"\U0001f780-\U0001f7ff\U0001f800-\U0001f8ff\U0001f900-\U0001f9ff\U0001fa00-\U0001fa6f"
"\U0001fa70-\U0001faffβ˜€-β›Ώβœ€-➿\U0001f1e6-\U0001f1ff]+",
flags=re.UNICODE,
)
_REPL = {"–": "-", "‑": "-", "β€”": "-", "_": " ", "β€œ": '"', "”": '"', "β€˜": "'", "’": "'",
"Β΄": "'", "`": "'", "[": " ", "]": " ", "|": " ", "/": " ", "#": " ", "β†’": " ", "←": " "}
_EXPR = {"@": " at ", "e.g.,": "for example, ", "i.e.,": "that is, "}
def preprocess_text(text: str, lang: str) -> str:
text = normalize("NFKD", text)
text = _EMOJI.sub("", text)
for k, v in _REPL.items():
text = text.replace(k, v)
text = re.sub(r"[β™₯β˜†β™‘Β©\\]", "", text)
for k, v in _EXPR.items():
text = text.replace(k, v)
for p in [",", r"\.", "!", r"\?", ";", ":", "'"]:
text = re.sub(r" " + p, p.replace("\\", ""), text)
while '""' in text:
text = text.replace('""', '"')
while "''" in text:
text = text.replace("''", "'")
while "``" in text:
text = text.replace("``", "`")
text = re.sub(r"\s+", " ", text).strip()
if not re.search(r"[.!?;:,'\"')\]}…。」』】〉》›»]$", text):
text += "."
if lang not in AVAILABLE_LANGS:
raise ValueError(f"Invalid language: {lang}")
return f"<{lang}>" + text + f"</{lang}>"
class TextProcessor:
def __init__(self, unicode_indexer_path: str):
self.indexer_path = unicode_indexer_path
with open(unicode_indexer_path) as f:
self.indexer = json.load(f) # list: code point -> id (-1 = unknown)
def encode(self, text: str, lang: str, preprocess: bool = True) -> list[int]:
s = preprocess_text(text, lang) if preprocess else text
return [self.indexer[ord(c)] if ord(c) < len(self.indexer) else -1 for c in s]
def unknown_chars(self, text: str, lang: str) -> set[str]:
s = preprocess_text(text, lang)
return {c for c in s if ord(c) >= len(self.indexer) or self.indexer[ord(c)] < 0}
def batch_numpy(self, texts, langs):
"""ids (B, T) int64 and mask (B, 1, T) float32 as numpy arrays."""
import numpy as np
seqs = [self.encode(t, l) for t, l in zip(texts, langs)]
T = max(len(s) for s in seqs)
ids = np.zeros((len(seqs), T), np.int64)
mask = np.zeros((len(seqs), 1, T), np.float32)
for i, s in enumerate(seqs):
ids[i, : len(s)] = s
mask[i, :, : len(s)] = 1
return ids, mask
def batch(self, texts: list[str], langs: list[str], device=None):
import torch
seqs = [self.encode(t, l) for t, l in zip(texts, langs)]
lengths = torch.tensor([len(s) for s in seqs])
ids = torch.zeros(len(seqs), int(lengths.max()), dtype=torch.long)
for i, s in enumerate(seqs):
ids[i, : len(s)] = torch.tensor(s)
mask = (torch.arange(ids.shape[1])[None] < lengths[:, None]).float().unsqueeze(1)
return ids.to(device), mask.to(device)