rita-cohere commited on
Commit
af3219f
·
verified ·
1 Parent(s): ff2ee71

Revert script.py to Julia original (parser v1 hurt under temp=0.8 resample)

Browse files
Files changed (1) hide show
  1. script.py +65 -140
script.py CHANGED
@@ -151,176 +151,101 @@ def expected_answer_count(query: str, task_type: str) -> int:
151
  return len(numbered) or 1
152
 
153
 
154
- def _letter_tokens(text: str) -> list[str] | None:
155
- """If text is only option letters (A J L … or A,J,L), return them uppercased."""
156
- parts = [p.strip("().[]") for p in re.split(r"[\s,;]+", text.strip()) if p.strip()]
157
- if len(parts) >= 2 and all(re.fullmatch(r"[A-Za-z]", p) for p in parts):
158
- return [p.upper() for p in parts]
159
- letters = re.findall(r"[A-Za-z]", text)
160
- if len(letters) >= 2 and re.fullmatch(r"[A-Za-z\s,;]+", text.strip()):
161
- return [c.upper() for c in letters]
162
- return None
163
-
164
-
165
  def split_single_line_answer(text: str, expected: int, task_type: str) -> list[str]:
166
  text = text.strip()
167
- if not text:
168
- return []
169
 
170
- def try_split(pattern: str, require_expected: bool) -> list[str] | None:
171
  parts = [part.strip() for part in re.split(pattern, text) if part.strip()]
172
- if require_expected:
173
- return parts if len(parts) == expected else None
174
- return parts if len(parts) >= 2 else None
175
 
176
  if task_type == "match_letters":
177
- letters = _letter_tokens(text)
178
- if letters is not None:
179
- if expected <= 1 or len(letters) == expected:
180
- return letters
181
- if len(letters) > expected:
182
- return letters[:expected]
183
- return letters
184
  for pattern in (r"\s+", r",\s*", r";\s*"):
185
- if result := try_split(pattern, require_expected=expected > 1):
186
- return [p.upper() if re.fullmatch(r"[A-Za-z]", p) else p for p in result]
 
 
 
187
  return [text]
188
 
189
- # Prefer exact expected count when known; else accept any multi-way split.
190
- require = expected > 1
191
  if task_type in ("text_to_num", "num_to_text"):
192
- patterns = (r",\s*", r";\s*", r"\s+")
193
- else:
194
- patterns = (r";\s*", r",\s*")
195
- for pattern in patterns:
196
- if result := try_split(pattern, require_expected=require):
197
- if expected > 1 and len(result) > expected:
198
- return result[:expected]
199
  return result
200
- # Last resort: comma/semicolon split even when expected heuristic said 1
201
- # (common failure: one line with many comma-joined answers).
202
- if text.count(",") >= 2 or text.count(";") >= 2:
203
- for pattern in (r";\s*", r",\s*"):
204
- parts = [p.strip() for p in re.split(pattern, text) if p.strip()]
205
- if len(parts) >= 3:
206
- return parts[:expected] if expected > 1 else parts
207
  return [text]
208
 
209
 
210
- def strip_bare_answer(ans: str, task_type: str) -> str:
211
- """Remove letter-label prefixes and English gloss tails; keep bare answer forms."""
212
- ans = ans.strip()
213
- if not ans:
214
- return ans
215
-
216
- if task_type == "match_letters":
217
- if re.fullmatch(r"[A-Za-z]", ans):
218
- return ans.upper()
219
- m = re.match(
220
- r"^\s*(?:\(([A-Za-z])\)|\[([A-Za-z])\]|([A-Za-z]))\s*[.):\-–—]?\s*(.*)$",
221
- ans,
222
- )
223
- if m:
224
- letter = m.group(1) or m.group(2) or m.group(3)
225
- rest = (m.group(4) or "").strip()
226
- if not rest or len(rest) <= 2:
227
- return letter.upper()
228
- letters = _letter_tokens(ans)
229
- if letters is not None and len(letters) == 1:
230
- return letters[0]
231
- return ans
232
-
233
- # "A: Luzia is my husband." / "B) foo" / "(C) bar"
234
- m = re.match(
235
- r"^\s*(?:\(([A-Za-z])\)|\[([A-Za-z])\]|([A-Za-z]))\s*[.):\-–—]\s+(.+)$",
236
- ans,
237
- )
238
- if m:
239
- ans = m.group(4).strip()
240
-
241
- # "form - to be called" / "form – English gloss"
242
- # Require an English gloss cue after the dash so language-internal hyphens stay.
243
- gloss = re.match(
244
- r"^(.+?)\s+[-–—]\s+((?:to|the|a|an|in|of|for|being)\b.*)$",
245
- ans,
246
- flags=re.IGNORECASE,
247
- )
248
- if gloss:
249
- rhs = gloss.group(2)
250
- non_ascii = sum(1 for c in rhs if ord(c) > 127)
251
- # Drop gloss if RHS is English-looking (few non-ASCII letters).
252
- if non_ascii <= 1:
253
- ans = gloss.group(1).strip()
254
-
255
- return ans
256
-
257
-
258
- def repair_answer_cardinality(answers: list[str], expected: int, task_type: str) -> list[str]:
259
- """Expand collapsed multi-answers; prefer matching expected count when known."""
260
- if not answers:
261
- return answers
262
-
263
- # Expand any match_letters line that is still a letter cluster.
264
- if task_type == "match_letters":
265
- expanded: list[str] = []
266
- for ans in answers:
267
- letters = _letter_tokens(ans)
268
- if letters is not None:
269
- expanded.extend(letters)
270
- elif re.fullmatch(r"[A-Za-z]", ans.strip()):
271
- expanded.append(ans.strip().upper())
272
- else:
273
- one = strip_bare_answer(ans, task_type)
274
- letters = _letter_tokens(one)
275
- expanded.extend(letters if letters is not None else [one])
276
- answers = expanded
277
-
278
- if len(answers) == 1:
279
- answers = split_single_line_answer(answers[0], expected, task_type)
280
-
281
- if expected > 1 and len(answers) > expected:
282
- answers = answers[:expected]
283
- return answers
284
-
285
-
286
  def postprocess_answer(text, query, task_type):
287
- """Keep only the lines after the last 'FINAL ANSWERS:' marker, one answer per line.
288
- Point 1: repair cardinality (letter/comma collapses). Point 2: bare-answer cleanup."""
 
289
  marker_match = list(re.finditer(r"(?im)^[^\w\n]*final answers?[^\w\n]*:?\s*$", text))
290
  if marker_match:
291
  text_after_marker = text[marker_match[-1].end():]
 
292
  else:
 
293
  return []
294
 
295
  answers = []
296
- for line in text_after_marker.splitlines():
297
- stripped_line = line.strip("`").strip()
298
- if stripped_line == "":
299
- continue
300
 
301
- match_numbered_prefix = re.match(r"^\s*\d+[.)]\s+(.*)", stripped_line)
302
- cleaned_line = (
303
- match_numbered_prefix.group(1).strip()
304
- if match_numbered_prefix
305
- else stripped_line
306
- )
307
- cleaned_line = re.sub(r"\*\*", "", cleaned_line).strip()
308
 
309
- if task_type == "match_letters":
310
- letters = _letter_tokens(cleaned_line)
311
- if letters is not None:
312
- answers.extend(letters)
313
- continue
314
- cleaned_line = strip_bare_answer(cleaned_line, task_type)
315
  else:
316
- cleaned_line = strip_bare_answer(cleaned_line, task_type)
317
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
318
  if cleaned_line:
319
  answers.append(cleaned_line)
320
 
 
 
 
 
321
  expected = expected_answer_count(query, task_type)
322
- answers = repair_answer_cardinality(answers, expected, task_type)
323
- answers = [strip_bare_answer(a, task_type) for a in answers if a.strip()]
 
 
324
  return answers
325
 
326
  rows = []
 
151
  return len(numbered) or 1
152
 
153
 
 
 
 
 
 
 
 
 
 
 
 
154
  def split_single_line_answer(text: str, expected: int, task_type: str) -> list[str]:
155
  text = text.strip()
156
+ if expected <= 1:
157
+ return [text]
158
 
159
+ def try_split(pattern: str) -> list[str] | None:
160
  parts = [part.strip() for part in re.split(pattern, text) if part.strip()]
161
+ return parts if len(parts) == expected else None
 
 
162
 
163
  if task_type == "match_letters":
 
 
 
 
 
 
 
164
  for pattern in (r"\s+", r",\s*", r";\s*"):
165
+ if result := try_split(pattern):
166
+ return result
167
+ letters = re.findall(r"[A-Za-z]", text)
168
+ if len(letters) == expected:
169
+ return [letter.upper() for letter in letters]
170
  return [text]
171
 
 
 
172
  if task_type in ("text_to_num", "num_to_text"):
173
+ for pattern in (r",\s*", r";\s*", r"\s+"):
174
+ if result := try_split(pattern):
175
+ return result
176
+ return [text]
177
+
178
+ for pattern in (r";\s*", r",\s*"):
179
+ if result := try_split(pattern):
180
  return result
 
 
 
 
 
 
 
181
  return [text]
182
 
183
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
184
  def postprocess_answer(text, query, task_type):
185
+ """Keep only the lines after the last 'FINAL ANSWERS:' marker, one answer per line,
186
+ stopping at the first empty line. If eval_type is multiple and only one line as answer, split at whitespace."""
187
+ # Updated regex to be more flexible with surrounding characters
188
  marker_match = list(re.finditer(r"(?im)^[^\w\n]*final answers?[^\w\n]*:?\s*$", text))
189
  if marker_match:
190
  text_after_marker = text[marker_match[-1].end():]
191
+ #print('FOUND FINAL ANSWER', text_after_marker)
192
  else:
193
+ #print("No 'FINAL ANSWERS:' marker found")
194
  return []
195
 
196
  answers = []
197
+ answer_lines = text_after_marker.splitlines()
198
+ for i, line in enumerate(answer_lines):
199
+ stripped_line = line.strip('`').strip()
 
200
 
201
+ # Stop processing if an empty line is encountered (not as first line)
202
+ if stripped_line=='':
203
+ continue
 
 
 
 
204
 
205
+ # Use a more precise regex to only remove numbering if it's a prefix to other text
206
+ # This ensures that lines which are just numbers (e.g., '1') are not stripped.
207
+ match_numbered_prefix = re.match(r"^\s*\d+[.)]\s+(.*)", stripped_line)
208
+ if match_numbered_prefix:
209
+ cleaned_line = match_numbered_prefix.group(1).strip()
 
210
  else:
211
+ cleaned_line = stripped_line
212
 
213
+ # Remove any bold markdown '**'
214
+ cleaned_line = re.sub(r"\*\*", "", cleaned_line).strip()
215
+
216
+ # Specific handling for 'match_letters' task type to strip extra words
217
+ if task_type == 'match_letters':
218
+ parts = [
219
+ part.strip("().[]")
220
+ for part in re.split(r"[\s,;]+", cleaned_line)
221
+ if part.strip()
222
+ ]
223
+ if not (len(parts) > 1 and all(re.fullmatch(r"[A-Za-z]", part) for part in parts)):
224
+ match_letter_word = re.match(
225
+ r"^\s*(?:\(([A-Za-z])\)|\[([A-Za-z])\]|([A-Za-z]))\.?:?\s*(.*)$",
226
+ cleaned_line,
227
+ )
228
+ if match_letter_word:
229
+ letter = (
230
+ match_letter_word.group(1)
231
+ or match_letter_word.group(2)
232
+ or match_letter_word.group(3)
233
+ )
234
+ cleaned_line = letter.upper()
235
+
236
+ # Append the cleaned, non-empty line
237
  if cleaned_line:
238
  answers.append(cleaned_line)
239
 
240
+ #print('PARSED ANSWERS', answers)
241
+
242
+ # Compare against QUERY length: sometimes model forgets newlines
243
+ #print('QUERY', query)
244
  expected = expected_answer_count(query, task_type)
245
+ query_len = len(query.splitlines()) - 2
246
+ #print(query_len)
247
+ if len(answers) == 1 and expected > 1:
248
+ answers = split_single_line_answer(answers[0], expected, task_type)
249
  return answers
250
 
251
  rows = []