"""Text frontend: NFKD normalisation, light cleanup, "text" 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"" 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)