Bina-0.2-Rizeh / pipeline /decode_bina02_persian.py
Reza2kn's picture
Publish Bina 0.2 Rizeh Stage-8 release
feef108 verified
Raw
History Blame Contribute Delete
10.5 kB
#!/usr/bin/env python3
from __future__ import annotations
import argparse
import json
import math
import re
import sqlite3
import unicodedata
from collections import defaultdict
from functools import lru_cache
from pathlib import Path
from typing import Any
from rapidfuzz.distance import Levenshtein
TOKEN_RE = re.compile(r"[\u0600-\u06ff]+")
DIACRITICS_RE = re.compile(r"[\u064b-\u065f\u0670\u06d6-\u06ed]")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--database", type=Path, required=True)
parser.add_argument("--input", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--details", type=Path)
parser.add_argument("--minimum-frequency", type=int, default=2)
parser.add_argument("--frequent-without-context", type=int, default=250)
parser.add_argument("--minimum-margin", type=float, default=4.0)
return parser.parse_args()
def normalize_word(value: str) -> str:
word = unicodedata.normalize("NFKC", value or "")
for source, target in (
("ي", "ی"),
("ك", "ک"),
("ۀ", "ه"),
("ة", "ه"),
("أ", "ا"),
("إ", "ا"),
("ٱ", "ا"),
("آ", "ا"),
("ؤ", "و"),
("ئ", "ی"),
("ى", "ی"),
):
word = word.replace(source, target)
word = DIACRITICS_RE.sub("", word)
return (
word.replace("\u200c", "")
.replace("\u200d", "")
.replace("\u200b", "")
)
def one_deletes(word: str) -> set[str]:
return {
word[:index] + word[index + 1 :]
for index in range(len(word))
}
class PersianLexiconDecoder:
def __init__(
self,
database: Path,
*,
minimum_frequency: int,
frequent_without_context: int,
minimum_margin: float,
) -> None:
uri = f"file:{database.resolve()}?mode=ro"
self.connection = sqlite3.connect(uri, uri=True)
self.minimum_frequency = minimum_frequency
self.frequent_without_context = frequent_without_context
self.minimum_margin = minimum_margin
self.words: dict[str, tuple[int, bool]] = {}
for word, frequency, hardword in self.connection.execute(
"SELECT word, frequency, hardword FROM unigram "
"WHERE frequency >= ? OR hardword = 1",
(minimum_frequency,),
):
self.words[str(word)] = (int(frequency), bool(hardword))
self.delete_index: dict[str, list[str]] = defaultdict(list)
for word in self.words:
if not 3 <= len(word) <= 32:
continue
for deletion in one_deletes(word):
self.delete_index[deletion].append(word)
@lru_cache(maxsize=500_000)
def bigram_frequency(self, left: str, right: str) -> int:
row = self.connection.execute(
"SELECT frequency FROM bigram "
"WHERE left_word = ? AND right_word = ?",
(left, right),
).fetchone()
return int(row[0]) if row else 0
@lru_cache(maxsize=500_000)
def candidates(self, word: str) -> tuple[str, ...]:
normalized = normalize_word(word)
possible: set[str] = set()
if normalized in self.words:
possible.add(normalized)
deletes = one_deletes(normalized)
for deletion in deletes:
if deletion in self.words:
possible.add(deletion)
for key in {normalized, *deletes}:
possible.update(self.delete_index.get(key, ()))
return tuple(possible)
def correction(
self,
word: str,
left: str | None,
right: str | None,
) -> tuple[str, dict[str, Any] | None]:
normalized = normalize_word(word)
if len(normalized) < 3 or normalized in self.words:
return word, None
maximum_distance = 1 if len(normalized) <= 6 else 2
ranked: list[tuple[float, str, int, int, int, bool]] = []
for candidate in self.candidates(normalized):
distance = Levenshtein.distance(
normalized,
candidate,
score_cutoff=maximum_distance,
)
if distance > maximum_distance:
continue
frequency, hardword = self.words[candidate]
left_frequency = (
self.bigram_frequency(left, candidate) if left else 0
)
right_frequency = (
self.bigram_frequency(candidate, right) if right else 0
)
contextual = left_frequency + right_frequency
if distance == 2 and not (
contextual
or frequency >= 2 * self.frequent_without_context
):
continue
if not contextual and (
frequency < self.frequent_without_context
):
continue
score = (
0.85 * math.log1p(frequency)
+ 1.35 * math.log1p(left_frequency)
+ 1.35 * math.log1p(right_frequency)
+ (0.8 if hardword else 0.0)
- 3.8 * distance
)
ranked.append(
(
score,
candidate,
distance,
frequency,
contextual,
hardword,
)
)
if not ranked:
return word, None
ranked.sort(reverse=True)
best = ranked[0]
second_score = ranked[1][0] if len(ranked) > 1 else -math.inf
margin = best[0] - second_score
if margin < self.minimum_margin:
return word, None
replacement = best[1]
return replacement, {
"original": word,
"normalized": normalized,
"replacement": replacement,
"distance": best[2],
"frequency": best[3],
"context_frequency": best[4],
"hardword": best[5],
"margin": round(margin, 6)
if math.isfinite(margin)
else None,
}
def decode(self, text: str) -> tuple[str, list[dict[str, Any]]]:
matches = list(TOKEN_RE.finditer(text))
if not matches:
return text, []
normalized_words = [
normalize_word(match.group(0)) for match in matches
]
replacements: list[str] = []
changes: list[dict[str, Any]] = []
for index, match in enumerate(matches):
left = normalized_words[index - 1] if index else None
right = (
normalized_words[index + 1]
if index + 1 < len(normalized_words)
else None
)
replacement, detail = self.correction(
match.group(0),
left,
right,
)
replacements.append(replacement)
if detail is not None:
detail["token_index"] = index
changes.append(detail)
output: list[str] = []
cursor = 0
for match, replacement in zip(matches, replacements):
output.append(text[cursor : match.start()])
output.append(replacement)
cursor = match.end()
output.append(text[cursor:])
return "".join(output), changes
def main() -> None:
args = parse_args()
decoder = PersianLexiconDecoder(
args.database,
minimum_frequency=args.minimum_frequency,
frequent_without_context=args.frequent_without_context,
minimum_margin=args.minimum_margin,
)
if args.output.exists():
raise FileExistsError(
f"Refusing to replace existing output: {args.output}"
)
if args.details and args.details.exists():
raise FileExistsError(
f"Refusing to replace existing details: {args.details}"
)
args.output.parent.mkdir(parents=True, exist_ok=True)
details_handle = (
args.details.open("w", encoding="utf-8")
if args.details
else None
)
rows = 0
changed_rows = 0
changed_tokens = 0
try:
with args.input.open(encoding="utf-8") as source:
with args.output.open("w", encoding="utf-8") as destination:
for line in source:
if not line.strip():
continue
row = json.loads(line)
prediction = str(
row.get(
"prediction",
row.get("hyp", row.get("hypothesis", "")),
)
or ""
)
decoded, changes = decoder.decode(prediction)
row["prediction_before_persian_decoder"] = prediction
row["prediction"] = decoded
row["persian_decoder_changes"] = len(changes)
destination.write(
json.dumps(row, ensure_ascii=False) + "\n"
)
rows += 1
if changes:
changed_rows += 1
changed_tokens += len(changes)
if details_handle is not None:
details_handle.write(
json.dumps(
{
"id": row.get("id"),
"split": row.get("split"),
"changes": changes,
},
ensure_ascii=False,
)
+ "\n"
)
finally:
if details_handle is not None:
details_handle.close()
decoder.connection.close()
print(
json.dumps(
{
"rows": rows,
"changed_rows": changed_rows,
"changed_tokens": changed_tokens,
"lexicon_words": len(decoder.words),
"delete_keys": len(decoder.delete_index),
},
ensure_ascii=False,
indent=2,
)
)
if __name__ == "__main__":
main()