tya-m1-temp06-user / script.py
rita-cohere's picture
M1 /think + Offelia tech: exact-N, phon, parser, targeted SC, induct
e41a455 verified
Raw
History Blame Contribute Delete
14.2 kB
"""IOL-AI 2026 — M1 (/think) + Offelia-style techniques.
Keep Tiny Aya reasoning (/think). Add:
- cardinality: count items, tell model exact N, truncate/pad
- task-aware + phonetic-bracket detector (Offelia)
- parser hygiene: drop essay lines after FINAL ANSWERS
- targeted self-consistency only on match_letters / fill_blanks (k=3)
- induction → apply (rules sheet then answers)
"""
import os
import subprocess
import sys
def _install_bundled_deps() -> None:
wheels_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "wheels")
if not os.path.isdir(wheels_dir):
return
subprocess.run(
[
sys.executable,
"-m",
"pip",
"install",
"-q",
"--no-index",
f"--find-links={wheels_dir}",
"transformers==4.56.2",
],
check=True,
)
_install_bundled_deps()
os.environ["HF_HUB_OFFLINE"] = "1"
os.environ["TRANSFORMERS_OFFLINE"] = "1"
MODEL_ID = "."
USER_THINK_TOKEN = "/think"
import json
import random
import re
from collections import Counter
import pandas as pd
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
END_THINKING = "<|END_THINKING|>"
START_THINKING = "<|START_THINKING|>"
THINKING_BUDGET = 1536
ANSWER_CONTINUATION_TOKENS = 512
COT_MAX_NEW_TOKENS = 1024
INDUCT_MAX_NEW_TOKENS = 512
THINK_TEMPERATURE = 0.6
THINK_TOP_P = 0.95
# Targeted SC only
SC_TASKS = frozenset({"match_letters", "fill_blanks"})
SC_K = 3
SYSTEM = "" # instructions on user turn (best M1 private recipe)
TASK_INSTRUCTIONS = {
"translation": (
"This is a TRANSLATION task. Give only the translated form, in the language "
"the task asks for. No explanation, no source form, just the translation."
),
"fill_blanks": (
"This is a FILL-IN-THE-BLANKS task. Give only the missing form for each blank, "
"nothing else."
),
"match_letters": (
"This is a MATCHING task. Each numbered item must be answered with a SINGLE "
"OPTION LETTER only (for example: C). Do NOT write the word, meaning, or "
"translation -- only the letter that matches."
),
"text_to_num": (
"This is a TEXT-TO-NUMBER task. Give the number in digits only (for example: 111)."
),
"num_to_text": (
"This is a NUMBER-TO-TEXT task. Write the number out in words, in the language "
"the task asks for. Give only the written-out form."
),
}
TASK_DEFAULT = (
"Give exactly what the instruction asks for, in the same form the examples use, "
"and nothing else."
)
PHONETIC_INSTRUCTION = (
"IMPORTANT -- this problem uses PHONETIC TRANSCRIPTION. The examples write forms "
"inside square brackets, like [bø:va]. Your answers must be phonetic transcriptions "
"in exactly that same notation: enclosed in square brackets, using the same phonetic "
"symbols. Do NOT give an English meaning or gloss -- give the transcribed FORM."
)
_IPA_HINT = re.compile(
r"[\u0250-\u02AF\u02B0-\u02FF\u0300-\u036F\u1D00-\u1D7Føœæðθŋɣʔ]"
)
_ASKS_NON_PHONETIC = re.compile(
r"(?i)translate\s+into\s+english"
r"|write\s+(it\s+)?in\s+the\s+[\w'\u2019-]+\s+orthography"
r"|in\s+the\s+regular\s+orthography"
)
_ASKS_TRANSCRIPTION = re.compile(r"(?i)\b(transcribe|transcription|phonetic(ally)?)\b")
_TURN_NOISE = re.compile(
r"<\|/?END_OF_TURN_TOKEN\|>|<\|/?START_OF_TURN_TOKEN\|>|"
r"<\|CHATBOT_TOKEN\|>|<EOS_TOKEN>|<BOS_TOKEN>"
)
_MARKER = re.compile(r"(?im)^\s*final answers?\s*:?\s*$")
def _bracketed_forms(text: str) -> list[str]:
out = []
for m in re.finditer(r"\[([^\[\]\n]{1,40})\]", text):
inner = m.group(1).strip()
if not inner or re.fullmatch(r"[\d\s,.\-]+", inner):
continue
out.append(inner)
return out
def is_phonetic_task(context: str, query: str, min_forms: int = 3) -> bool:
if _ASKS_NON_PHONETIC.search(query):
return False
if _bracketed_forms(query) and not _ASKS_TRANSCRIPTION.search(query):
return False
forms = _bracketed_forms(context) + _bracketed_forms(query)
if len(forms) < min_forms:
return False
phonetic_looking = sum(1 for f in forms if _IPA_HINT.search(f) or ":" in f)
return phonetic_looking >= max(2, len(forms) // 4)
def count_items(query: str) -> int:
n = len(re.findall(r"(?m)^\s*\d+[.)]", query))
if n:
return n
if "blanks" in query.lower():
m = re.search(r"\((\d+)-(\d+)\)", query)
if m:
return int(m.group(2)) - int(m.group(1)) + 1
return len(re.findall(r"\(\d+\)", query)) or 0
return 0
def _looks_like_prose(line: str) -> bool:
if re.search(
r"(?i)^(final answers?|answers?|note|reviewing|summary|explanation|verification)\b.*:$",
line,
):
return True
if re.search(
r"(?i)^(here (are|is)|the (final )?answers? (are|is)|based on|therefore|thus|"
r"in summary|colors? are expressed|these stems)\b",
line,
):
return True
if line.rstrip().endswith(":") and len(line) > 3:
return True
if len(line) > 120:
return True
return False
def _strip_gloss_keep_form(line: str) -> str:
s = re.sub(r"\*\*", "", (line or "").strip())
s = re.split(r"\s+_?(?:GCY|NS|N/A)_?\b", s, maxsplit=1, flags=re.I)[0].strip()
m = re.match(
r"^(.+?)\s+[-–—]\s+((?:to|the|a|an|in|of|for|being|means?)\b.*)$",
s,
flags=re.I,
)
if m:
s = m.group(1).strip()
return s.strip()
def parse_answers(text: str, n_items: int = 0) -> list[str]:
text = after_thinking(text)
markers = list(_MARKER.finditer(text))
if markers:
text = text[markers[-1].end() :]
answers = []
for line in text.splitlines():
line = re.sub(r"^\s*\d+[.)]\s*", "", line).strip().strip("`")
if not line or _looks_like_prose(line):
continue
line = _strip_gloss_keep_form(line)
if not line:
continue
# match_letters letter blob
if re.fullmatch(r"(?:[A-Za-z]\s+)+[A-Za-z]", line):
answers.extend([p.upper() for p in line.split()])
continue
answers.append(line)
if n_items > 0:
answers = answers[:n_items]
if len(answers) < n_items:
answers += [""] * (n_items - len(answers))
return answers
def after_thinking(text: str) -> str:
if END_THINKING in text:
text = text.rsplit(END_THINKING, 1)[-1]
elif START_THINKING in text:
text = ""
return _TURN_NOISE.sub("", text)
def build_instructions(task_type: str, context: str, query: str) -> str:
specific = TASK_INSTRUCTIONS.get(str(task_type).strip().lower(), TASK_DEFAULT)
parts = [
"You solve International Linguistics Olympiad (IOL) problems from the data you are given.",
specific,
"Put answers ONLY after a line that says exactly: FINAL ANSWERS:",
"Bare answers only: no numbering, no quotes, no commentary, no _GCY/_NS glosses.",
"Never dump the alphabet. Never write an essay under FINAL ANSWERS:.",
]
if is_phonetic_task(context, query):
parts.append(PHONETIC_INSTRUCTION)
return "\n\n".join(parts)
def build_user(
instructions: str,
context: str,
query: str,
*,
n_items: int,
think_token: str = "",
rules: str = "",
mode: str = "answer",
) -> str:
parts = [instructions.strip(), "", context.strip()]
if rules.strip():
parts += ["", "RULES:", rules.strip()]
parts += ["", query.strip()]
if mode == "induct":
parts += [
"",
"Deduce linguistic RULES from CONTEXT only. Do NOT answer QUERY.",
"Write a bullet list under a line that says exactly: RULES:",
]
elif n_items > 0:
parts += [
"",
f"There are exactly {n_items} items to answer. "
f"Give exactly {n_items} answers after FINAL ANSWERS:, "
"one per line, no more and no fewer.",
]
if think_token:
parts.append(think_token.strip())
return "\n".join(parts)
def _end_thinking_id(tok) -> int:
end_id = tok.convert_tokens_to_ids(END_THINKING)
if end_id is None or end_id == tok.unk_token_id:
ids = tok.encode(END_THINKING, add_special_tokens=False)
if len(ids) == 1:
end_id = ids[0]
if end_id is None or end_id == tok.unk_token_id:
raise RuntimeError(f"missing {END_THINKING}")
return int(end_id)
def _build_prompt_ids(tok, user: str, *, thinking: bool):
messages = [{"role": "user", "content": user}]
try:
return tok.apply_chat_template(
messages,
add_generation_prompt=True,
return_tensors="pt",
reasoning_options={"enabled": thinking},
)
except TypeError:
return tok.apply_chat_template(
messages, add_generation_prompt=True, return_tensors="pt"
)
@torch.inference_mode()
def generate_with_think(
model,
tok,
prompt_ids,
end_id: int,
*,
sample_think: bool,
think_budget: int = THINKING_BUDGET,
answer_tokens: int = ANSWER_CONTINUATION_TOKENS,
):
device = next(model.parameters()).device
prompt_ids = prompt_ids.to(device)
prompt_len = prompt_ids.shape[-1]
think_kw = (
dict(do_sample=True, temperature=THINK_TEMPERATURE, top_p=THINK_TOP_P)
if sample_think
else dict(do_sample=False)
)
think_out = model.generate(
prompt_ids,
max_new_tokens=think_budget,
pad_token_id=tok.pad_token_id or tok.eos_token_id,
**think_kw,
)[0]
gen_ids = think_out[prompt_len:].tolist()
if end_id not in gen_ids:
cont = torch.cat(
[think_out, torch.tensor([end_id], device=device, dtype=think_out.dtype)]
)
else:
cont = think_out
# greedy answer
full = model.generate(
cont.unsqueeze(0),
max_new_tokens=answer_tokens,
do_sample=False,
pad_token_id=tok.pad_token_id or tok.eos_token_id,
)[0]
return _TURN_NOISE.sub("", tok.decode(full[prompt_len:], skip_special_tokens=False)).strip()
@torch.inference_mode()
def generate_plain(model, tok, prompt_ids, max_new: int, *, sample: bool = False):
device = next(model.parameters()).device
prompt_ids = prompt_ids.to(device)
prompt_len = prompt_ids.shape[-1]
kw = (
dict(do_sample=True, temperature=0.6, top_p=0.95)
if sample
else dict(do_sample=False)
)
out = model.generate(
prompt_ids,
max_new_tokens=max_new,
pad_token_id=tok.pad_token_id or tok.eos_token_id,
**kw,
)[0]
return _TURN_NOISE.sub("", tok.decode(out[prompt_len:], skip_special_tokens=False)).strip()
def majority_vote(samples: list[list[str]], n_items: int) -> list[str]:
usable = [s for s in samples if any(x.strip() for x in s)]
if not usable:
return [""] * max(n_items, 0)
n = n_items or max(len(s) for s in usable)
padded = [(list(s) + [""] * n)[:n] for s in usable]
# prefer full-tuple agreement
counts = Counter(tuple(p) for p in padded)
best, c = counts.most_common(1)[0]
if c >= 2:
return list(best)
return [Counter(p[i] for p in padded).most_common(1)[0][0] for i in range(n)]
def extract_rules(text: str) -> str:
text = after_thinking(text)
m = list(re.finditer(r"(?im)^\s*rules?\s*:?\s*$", text))
if m:
return text[m[-1].end() :].strip()[:2000]
return text.strip()[:2000]
tok = AutoTokenizer.from_pretrained(MODEL_ID)
end_id = _end_thinking_id(tok)
model = AutoModelForCausalLM.from_pretrained(
MODEL_ID, torch_dtype=torch.float16, device_map="auto"
).eval()
df = pd.read_csv("/tmp/data/test.csv", dtype=str).fillna("")
rows = []
for i, r in df.iterrows():
task = str(r.get("task_type", "") or "")
n_items = count_items(r["query"])
instr = build_instructions(task, r["context"], r["query"])
# Pass A: induction (greedy think)
induct_user = build_user(
instr,
r["context"],
r["query"],
n_items=0,
think_token=USER_THINK_TOKEN,
mode="induct",
)
induct_ids = _build_prompt_ids(tok, induct_user, thinking=True)
# Short greedy think for rules only — keep T4 headroom for answer (+ SC).
induct_text = generate_with_think(
model,
tok,
induct_ids,
end_id,
sample_think=False,
think_budget=INDUCT_MAX_NEW_TOKENS,
answer_tokens=256,
)
rules = extract_rules(induct_text)
# Pass B: answer with rules
def one_answer(seed: int, sample_think: bool) -> list[str]:
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
user = build_user(
instr,
r["context"],
r["query"],
n_items=n_items,
think_token=USER_THINK_TOKEN,
rules=rules,
mode="answer",
)
ids = _build_prompt_ids(tok, user, thinking=True)
text = generate_with_think(
model, tok, ids, end_id, sample_think=sample_think
)
return parse_answers(text, n_items=n_items)
if task in SC_TASKS:
samples = [
one_answer(1000 + int(i) * 97 + k * 17, sample_think=True)
for k in range(SC_K)
]
answers = majority_vote(samples, n_items)
print(f" targeted SC k={SC_K} task={task}", flush=True)
else:
answers = one_answer(1000 + int(i) * 97, sample_think=True)
# Fallback: no-rules single greedy think if mostly empty
if n_items > 0 and sum(1 for a in answers if a.strip()) < max(1, n_items // 2):
answers = one_answer(42 + int(i), sample_think=False)
rows.append({"id": r["id"], "pred": json.dumps(answers, ensure_ascii=False)})
pd.DataFrame(rows).to_csv("submission.csv", index=False)
print(f"[{i + 1}/{len(df)}] n={n_items} got={len(answers)} phon={is_phonetic_task(r['context'], r['query'])}", flush=True)
print("wrote submission.csv", flush=True)