""" Translate a stratified 150k-row subset of Ajhesh7/gsm8k-reasoning-SFT-datas into Arabic with ByteDance-Seed/Seed-X-PPO-7B (direct, per-row — numbers and names kept in place). Unique segments are translated once (greedy decoding is deterministic, so identical input gives identical output) and cached to a jsonl that makes the run resumable. Usage: python translate_gsm.py [n_rows] """ import json import os import random import re import sys import time from collections import defaultdict import pyarrow.parquet as pq sys.path.insert(0, ".") from gsm_common import NUM_RE, SRC_PROMPT, parse SRC_PARQUET = "data_gsm/data/train-00000-of-00001.parquet" CACHE = "out_gsm/translations.jsonl" N_ROWS = int(sys.argv[1]) if len(sys.argv) > 1 else 150_000 CHUNK = 20_000 SEED = 42 MODEL = "./models/Seed-X-PPO-7B" NAME_RE = re.compile(r"\b[A-Z][a-z]{2,}\b") def select_rows(texts, n, seed=SEED): """Stratify by name/number-agnostic question pattern so every pattern is represented.""" rows = [] for i, t in enumerate(texts): p = parse(t) if p: rows.append((i,) + p) buckets = defaultdict(list) for r in rows: key = NAME_RE.sub("@", NUM_RE.sub("#", r[1])) buckets[key].append(r) rng = random.Random(seed) for b in buckets.values(): rng.shuffle(b) chosen, leftover = [], [] floor = max(1, n // (len(buckets) * 4)) for b in buckets.values(): chosen.extend(b[:floor]) leftover.extend(b[floor:]) rng.shuffle(leftover) chosen.extend(leftover[: max(0, n - len(chosen))]) rng.shuffle(chosen) print(f"[*] {len(buckets)} question patterns; floor {floor}/pattern; selected {len(chosen)} rows") return chosen[:n] def load_cache(): done = {} if os.path.exists(CACHE): with open(CACHE, encoding="utf-8") as fh: for line in fh: try: rec = json.loads(line) done[rec["src"]] = rec["tgt"] except json.JSONDecodeError: continue # truncated last line from a killed run return done def main(): os.makedirs("out_gsm", exist_ok=True) texts = pq.read_table(SRC_PARQUET).to_pydict()["text"] rows = select_rows(texts, N_ROWS) with open("out_gsm/selected_rows.jsonl", "w", encoding="utf-8") as fh: for idx, q, t, a in rows: fh.write(json.dumps({"idx": idx, "question": q, "thinking": t, "answer": a}, ensure_ascii=False) + "\n") segments = [] seen = set() for _, q, t, _ in rows: for s in (q, t): if s not in seen: seen.add(s) segments.append(s) done = load_cache() todo = [s for s in segments if s not in done] print(f"[*] {len(rows)} rows -> {len(segments)} unique segments; {len(done)} cached, {len(todo)} to translate") if not todo: print("[+] nothing to do") return from vllm import LLM, SamplingParams llm = LLM(model=MODEL, max_num_seqs=512, gpu_memory_utilization=0.92, max_model_len=1024) params = SamplingParams(temperature=0, max_tokens=256, skip_special_tokens=True) start = time.time() with open(CACHE, "a", encoding="utf-8") as fh: for i in range(0, len(todo), CHUNK): chunk = todo[i : i + CHUNK] outs = llm.generate([SRC_PROMPT.format(text=s) for s in chunk], params) for src, o in zip(chunk, outs): fh.write(json.dumps({"src": src, "tgt": o.outputs[0].text.strip()}, ensure_ascii=False) + "\n") fh.flush() done_n = i + len(chunk) rate = done_n / (time.time() - start) eta = (len(todo) - done_n) / rate / 60 print(f"[*] {done_n}/{len(todo)} segments {rate:.1f} seg/s ETA {eta:.0f} min", flush=True) print(f"[+] done in {(time.time() - start)/60:.1f} min -> {CACHE}") if __name__ == "__main__": main()