import torch
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
repo = "olaverse/diacnet-mini-2.0"
tok = AutoTokenizer.from_pretrained(repo)
model = AutoModelForSeq2SeqLM.from_pretrained(repo, dtype=torch.bfloat16).eval() # float32 on CPU
import difflib
import unicodedata as ud
_FOLD = str.maketrans("ɓɗƙƴđıłƁƊƘƳĐŁ", "bdkydilBDKYDL") # letters with no combining form
_LETTER = {"\u0653", "\u0654", "\u0655"} # Arabic madda / hamza are spelling, not marks
def _units(text):
units = [] # [letter, letter + marks]
for c in ud.normalize("NFD", text):
if units and ud.combining(c):
if c in _LETTER:
units[-1][0] += c
units[-1][1] += c
else:
units.append([c, c])
return [(ud.normalize("NFC", b).translate(_FOLD), ud.normalize("NFC", f)) for b, f in units]
def align(source, output):
"""`source` with the marks diacnet put on every letter it kept; letters it changed,
dropped or added fall back to `source`, so your text itself never changes."""
src, out = _units(source), _units(output)
res = [f for _, f in src]
sm = difflib.SequenceMatcher(None, [b for b, _ in src], [b for b, _ in out], autojunk=False)
# ... see the full example on the model card