Roni Egbu commited on
Commit
01aa95b
Β·
1 Parent(s): 1da4b8d

feat: implement context-aware pronoun resolution in translation process

Browse files
Files changed (1) hide show
  1. backend/app/models/mt_model.py +66 -64
backend/app/models/mt_model.py CHANGED
@@ -1,7 +1,9 @@
 
1
  import threading
2
  import time
3
  from transformers import MarianMTModel, MarianTokenizer
4
 
 
5
  MODEL_MAP: dict[tuple[str, str], str] = {
6
  ("en", "fr"): "Helsinki-NLP/opus-mt-en-fr",
7
  ("fr", "en"): "Helsinki-NLP/opus-mt-fr-en",
@@ -27,6 +29,38 @@ PIVOT_PAIRS: frozenset[tuple[str, str]] = frozenset({
27
 
28
  CONTEXT_WINDOW = 3
29
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
30
 
31
  class HelsinkiTranslator:
32
  def __init__(self):
@@ -47,16 +81,14 @@ class HelsinkiTranslator:
47
  def _load_pair(self, src: str, tgt: str) -> None:
48
  pair = (src, tgt)
49
  lock = self._get_load_lock(pair)
50
-
51
  with lock:
52
  if pair in self._models:
53
  return
54
-
55
  model_name = MODEL_MAP.get(pair)
56
  if model_name is None:
57
- raise ValueError(f"[MT] No direct model for {src}β†’{tgt}. "
58
- f"Use ensure_pair_loaded for pivot pairs.")
59
-
60
  print(f"[MT] Loading {src}β†’{tgt} ({model_name})...")
61
  t0 = time.time()
62
  tokenizer = MarianTokenizer.from_pretrained(model_name)
@@ -71,7 +103,6 @@ class HelsinkiTranslator:
71
  pair = (src, tgt)
72
  tokenizer = self._tokenizers[pair]
73
  model = self._models[pair]
74
-
75
  inputs = tokenizer(
76
  text,
77
  return_tensors="pt",
@@ -87,15 +118,13 @@ class HelsinkiTranslator:
87
 
88
  def ensure_pair_loaded(self, src: str, tgt: str) -> None:
89
  pair = (src, tgt)
90
-
91
  if pair in MODEL_MAP:
92
  if pair not in self._models:
93
  self._load_pair(src, tgt)
94
  elif pair in PIVOT_PAIRS:
95
- # Need src→en and en→tgt
96
- if (src, "en") in MODEL_MAP and (src, "en") not in self._models:
97
  self._load_pair(src, "en")
98
- if ("en", tgt) in MODEL_MAP and ("en", tgt) not in self._models:
99
  self._load_pair("en", tgt)
100
  else:
101
  print(f"[MT] ⚠️ No route defined for {src}β†’{tgt}")
@@ -111,60 +140,33 @@ class HelsinkiTranslator:
111
  if src == tgt:
112
  return text
113
 
114
- pair = (src, tgt)
115
-
116
- # ── Context-aware translation ─────────────────────────────────────────
117
-
118
- if use_context and context:
119
  recent = context[-CONTEXT_WINDOW:]
120
- context_str = " ".join(recent)
121
- full_str = context_str + " " + text
122
-
123
- if pair in MODEL_MAP:
124
- if pair not in self._models:
125
- self._load_pair(src, tgt)
126
- context_translated = self._translate_direct(context_str, src, tgt)
127
- full_translated = self._translate_direct(full_str, src, tgt)
128
-
129
- elif pair in PIVOT_PAIRS:
130
- en_pair = (src, "en")
131
- tgt_pair = ("en", tgt)
132
- if en_pair not in self._models: self._load_pair(src, "en")
133
- if tgt_pair not in self._models: self._load_pair("en", tgt)
134
-
135
- ctx_en = self._translate_direct(context_str, src, "en")
136
- full_en = self._translate_direct(full_str, src, "en")
137
- context_translated = self._translate_direct(ctx_en, "en", tgt)
138
- full_translated = self._translate_direct(full_en, "en", tgt)
139
-
140
- else:
141
- print(f"[MT] ⚠️ No route for {src}β†’{tgt}, returning original")
142
- return text
143
-
144
- prefix = context_translated.strip()
145
- result = full_translated.strip()
146
 
147
- if prefix and result.startswith(prefix):
148
- result = result[len(prefix):].lstrip(" ,.;")
149
-
150
- return result.strip() if result.strip() else full_translated.strip()
151
 
152
- # ── Sentence-level translation ────────────────────────────────────────────
153
- else:
154
- input_text = text
155
-
156
- if pair in MODEL_MAP:
157
- if pair not in self._models:
158
- self._load_pair(src, tgt)
159
- return self._translate_direct(input_text, src, tgt)
160
-
161
- if pair in PIVOT_PAIRS:
162
- en_pair = (src, "en")
163
- tgt_pair = ("en", tgt)
164
- if en_pair not in self._models: self._load_pair(src, "en")
165
- if tgt_pair not in self._models: self._load_pair("en", tgt)
166
- en_text = self._translate_direct(input_text, src, "en")
167
- return self._translate_direct(en_text, "en", tgt)
168
-
169
- print(f"[MT] ⚠️ No route for {src}β†’{tgt}, returning original")
170
- return text
 
 
 
 
1
+ import re
2
  import threading
3
  import time
4
  from transformers import MarianMTModel, MarianTokenizer
5
 
6
+ # ── Language pair β†’ HuggingFace model name ───────────────────────────────────
7
  MODEL_MAP: dict[tuple[str, str], str] = {
8
  ("en", "fr"): "Helsinki-NLP/opus-mt-en-fr",
9
  ("fr", "en"): "Helsinki-NLP/opus-mt-fr-en",
 
29
 
30
  CONTEXT_WINDOW = 3
31
 
32
+ _PRONOUN_SUBJECT = re.compile(
33
+ r'\b(he|she|they|it|him|her|them|his|hers|their|its)\b',
34
+ re.IGNORECASE
35
+ )
36
+
37
+ _NAME_PATTERN = re.compile(r'\b[A-Z][a-z]{2,}\b')
38
+
39
+
40
+ def _resolve_pronouns(text: str, context: list[str]) -> str:
41
+ if not _PRONOUN_SUBJECT.search(text):
42
+ return text
43
+
44
+ candidate = None
45
+ for utt in reversed(context):
46
+ names = _NAME_PATTERN.findall(utt)
47
+ names = [n for n in names if n not in (
48
+ "The", "A", "An", "This", "That", "These", "Those",
49
+ "I", "You", "We", "They", "He", "She", "It"
50
+ )]
51
+ if names:
52
+ candidate = names[-1]
53
+ break
54
+
55
+ if not candidate:
56
+ return text
57
+
58
+ resolved = _PRONOUN_SUBJECT.sub(candidate, text)
59
+ if resolved != text:
60
+ print(f"[MT] Pronoun resolved: '{text}' β†’ '{resolved}' "
61
+ f"(candidate: {candidate})")
62
+ return resolved
63
+
64
 
65
  class HelsinkiTranslator:
66
  def __init__(self):
 
81
  def _load_pair(self, src: str, tgt: str) -> None:
82
  pair = (src, tgt)
83
  lock = self._get_load_lock(pair)
 
84
  with lock:
85
  if pair in self._models:
86
  return
 
87
  model_name = MODEL_MAP.get(pair)
88
  if model_name is None:
89
+ raise ValueError(
90
+ f"[MT] No direct model for {src}β†’{tgt}."
91
+ )
92
  print(f"[MT] Loading {src}β†’{tgt} ({model_name})...")
93
  t0 = time.time()
94
  tokenizer = MarianTokenizer.from_pretrained(model_name)
 
103
  pair = (src, tgt)
104
  tokenizer = self._tokenizers[pair]
105
  model = self._models[pair]
 
106
  inputs = tokenizer(
107
  text,
108
  return_tensors="pt",
 
118
 
119
  def ensure_pair_loaded(self, src: str, tgt: str) -> None:
120
  pair = (src, tgt)
 
121
  if pair in MODEL_MAP:
122
  if pair not in self._models:
123
  self._load_pair(src, tgt)
124
  elif pair in PIVOT_PAIRS:
125
+ if (src, "en") not in self._models:
 
126
  self._load_pair(src, "en")
127
+ if ("en", tgt) not in self._models:
128
  self._load_pair("en", tgt)
129
  else:
130
  print(f"[MT] ⚠️ No route defined for {src}β†’{tgt}")
 
140
  if src == tgt:
141
  return text
142
 
143
+ # Apply pronoun resolution heuristic (English source only for now)
144
+ input_text = text
145
+ if use_context and context and src == "en":
 
 
146
  recent = context[-CONTEXT_WINDOW:]
147
+ input_text = _resolve_pronouns(text, recent)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
148
 
149
+ pair = (src, tgt)
 
 
 
150
 
151
+ # ── Direct translation ────────────────────────────────────────────────
152
+ if pair in MODEL_MAP:
153
+ if pair not in self._models:
154
+ self._load_pair(src, tgt)
155
+ result = self._translate_direct(input_text, src, tgt)
156
+ print(f"[MT] context={'on' if use_context and context else 'off'} "
157
+ f"input='{input_text[:60]}' output='{result[:60]}'")
158
+ return result
159
+
160
+ # ── Pivot through English ─────────────────────────────────────────────
161
+ if pair in PIVOT_PAIRS:
162
+ en_pair = (src, "en")
163
+ tgt_pair = ("en", tgt)
164
+ if en_pair not in self._models: self._load_pair(src, "en")
165
+ if tgt_pair not in self._models: self._load_pair("en", tgt)
166
+ en_text = self._translate_direct(input_text, src, "en")
167
+ result = self._translate_direct(en_text, "en", tgt)
168
+ print(f"[MT] pivot {src}→en→{tgt}: '{result[:60]}'")
169
+ return result
170
+
171
+ print(f"[MT] ⚠️ No route for {src}β†’{tgt}, returning original")
172
+ return text