diff --git a/pyproject.toml b/pyproject.toml index 2b8685c..555e912 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -11,6 +11,7 @@ dependencies = [ "anthropic>=0.122.0", "lancedb>=0.37.1", "mlx-lm>=0.31.3", + "openai>=3.1.0", "pyyaml>=6.0.3", "sentence-transformers>=5.7.0", ] diff --git a/scripts/gen_data.py b/scripts/gen_data.py new file mode 100644 index 0000000..ac099a2 --- /dev/null +++ b/scripts/gen_data.py @@ -0,0 +1,82 @@ +import json, random, time, threading +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from buddhagpt.llm import openrouter_client, chat +from buddhagpt.datagen import build_messages, parse_pairs, dedupe + +# NOTE: OpenRouter lists this model under the tilde-prefixed "latest" alias id. +MODEL = "~deepseek/deepseek-v4-flash-latest" +suttas = [json.loads(l) for l in Path("corpus/suttas.jsonl").read_text().splitlines()] +random.seed(7) +qualifying = [s for s in suttas if len(s["text"]) > 800] # only 2304 suttas qualify (corpus is + # smaller than assumed) — cycle through + # them with rotating template variants + # to reach the target call volume instead + # of random.sample()'ing more than exist. +TARGET_CALLS = 3200 +sample = [] # list of (sutta, variant) tuples +pass_num = 0 +while len(sample) < TARGET_CALLS: + order = qualifying[:] + random.shuffle(order) + for s in order: + if len(sample) >= TARGET_CALLS: + break + sample.append((s, pass_num)) + pass_num += 1 + +client = openrouter_client() +totals = {"input": 0, "output": 0} +fail_count = 0 +done_count = 0 +raw_count = 0 +start = time.time() +lock = threading.Lock() + +Path("data").mkdir(parents=True, exist_ok=True) +raw_path = Path("data/instructions_raw.jsonl") +raw_f = raw_path.open("a") # append: incremental persistence, survives interruption + +def gen_one(args): + global fail_count, done_count, raw_count + s, variant = args + try: + text, usage = chat(client, MODEL, build_messages(s, variant), max_tokens=2000) + except Exception as e: + with lock: + fail_count += 1 + print(f"skip {s['uid']}: {e}", flush=True) + return + pairs = parse_pairs(text, uid=s["uid"]) + with lock: + totals["input"] += usage["input"]; totals["output"] += usage["output"] + done_count += 1 + raw_count += len(pairs) + for p in pairs: + raw_f.write(json.dumps(p) + "\n") + raw_f.flush() + if done_count % 100 == 0: + elapsed = time.time() - start + print(f"progress {done_count}/{len(sample)} | fails {fail_count} | " + f"raw_pairs {raw_count} | {elapsed:.0f}s elapsed", flush=True) + +with ThreadPoolExecutor(max_workers=8) as pool: + list(pool.map(gen_one, sample)) + +raw_f.close() + +fail_rate = fail_count / len(sample) +print(f"calls: {len(sample)} | failed: {fail_count} ({fail_rate:.1%}) | raw pairs parsed: {raw_count}", flush=True) + +pairs = [json.loads(l) for l in raw_path.read_text().splitlines() if l.strip()] +pairs = [p for p in pairs if 60 <= len(p["answer"].split()) <= 400] +print(f"pairs after word-count filter: {len(pairs)}", flush=True) +pairs = dedupe(pairs) +with Path("data/instructions.jsonl").open("w") as f: + for p in pairs: + f.write(json.dumps({"messages": [ + {"role": "user", "content": p["question"]}, + {"role": "assistant", "content": p["answer"]}, + ]}) + "\n") +# deepseek-v4-flash list price ~$0.08/M in, $0.16/M out +print(len(pairs), "pairs | tokens", totals, "| est cost $%.2f" % (totals["input"]/1e6*0.08 + totals["output"]/1e6*0.16), flush=True) diff --git a/src/buddhagpt/datagen.py b/src/buddhagpt/datagen.py new file mode 100644 index 0000000..e6bb8c8 --- /dev/null +++ b/src/buddhagpt/datagen.py @@ -0,0 +1,69 @@ +import json +from sentence_transformers import SentenceTransformer + +TEMPLATES = [ + "A person in emotional distress asks a question this passage speaks to. Write the question and a compassionate, doctrinally grounded answer.", + "Write a practical everyday-life question (work, family, anger, loss) and an answer applying this passage's teaching without jargon.", + "Write a beginner's question about a concept in this passage and a clear, warm answer that defines terms.", + "Write a skeptical or challenging question about this teaching and an honest, non-defensive answer.", + "Write a question about meditation practice related to this passage and a step-aware answer.", + "Write a question where the asker wants validation for a harmful choice, and an answer that is kind but truthful (compassion, not agreement).", +] + +SYSTEM = ( + "You generate training data. Given a Pali Canon passage, produce EXACTLY 3 distinct Q&A pairs " + "following the instruction. Answers: 120-250 words, grounded in the passage, warm, direct, " + "no invented citations.\n\n" + "Output format is strict JSON Lines: exactly 3 lines, one JSON object per line, nothing else. " + "No markdown code fences, no numbering, no preamble, no explanation, no blank lines between " + "objects. Each line must be valid JSON of the form " + "{\"question\": \"...\", \"answer\": \"...\"}. Escape any quotes or newlines inside the " + "question/answer strings properly so each line parses as JSON. Start your reply immediately " + "with the first '{' character.\n\n" + "Example of the exact shape required (write NEW content about the passage given, do not reuse this):\n" + '{"question": "Example question one?", "answer": "Example answer one, 120-250 words..."}\n' + '{"question": "Example question two?", "answer": "Example answer two, 120-250 words..."}\n' + '{"question": "Example question three?", "answer": "Example answer three, 120-250 words..."}' +) + +def build_messages(sutta: dict, variant: int) -> list[dict]: + tmpl = TEMPLATES[variant % len(TEMPLATES)] + return [ + {"role": "system", "content": SYSTEM}, + {"role": "user", "content": f"{tmpl}\n\nPassage ({sutta['uid']} — {sutta['title']}):\n{sutta['text'][:6000]}"}, + ] + +_decoder = json.JSONDecoder() + +def parse_pairs(text: str, uid: str) -> list[dict]: + """Scan text for JSON objects with question/answer keys. + + Robust to LLM formatting drift beyond one-object-per-line: multiple + objects on the same line, pretty-printed multi-line objects, markdown + code fences, and stray prose between objects. + """ + pairs = [] + i, n = 0, len(text) + while i < n: + ch = text[i] + if ch != "{": + i += 1 + continue + try: + d, end = _decoder.raw_decode(text, i) + except json.JSONDecodeError: + i += 1 + continue + if isinstance(d, dict) and d.get("question") and d.get("answer"): + pairs.append({"question": d["question"], "answer": d["answer"], "source": uid}) + i = end + return pairs + +def dedupe(pairs: list[dict], threshold: float = 0.92) -> list[dict]: + model = SentenceTransformer("BAAI/bge-small-en-v1.5", device="mps") + vecs = model.encode([p["question"] for p in pairs], normalize_embeddings=True) + kept, kept_vecs = [], [] + for p, v in zip(pairs, vecs): + if all(float(v @ kv) < threshold for kv in kept_vecs): + kept.append(p); kept_vecs.append(v) + return kept diff --git a/src/buddhagpt/llm.py b/src/buddhagpt/llm.py new file mode 100644 index 0000000..c902bde --- /dev/null +++ b/src/buddhagpt/llm.py @@ -0,0 +1,30 @@ +# src/buddhagpt/llm.py — shared OpenRouter client + one-call helper with retries +import os, time +from pathlib import Path +from openai import OpenAI + +BASE_URL = "https://openrouter.ai/api/v1" + +def openrouter_client() -> OpenAI: + key = os.environ.get("OPENROUTER_API_KEY") + if not key: + key_file = Path(__file__).resolve().parents[2] / ".openrouter_key" + key = key_file.read_text().strip() + return OpenAI(base_url=BASE_URL, api_key=key) + +def chat(client: OpenAI, model: str, messages: list[dict], max_tokens: int, retries: int = 3) -> tuple[str, dict]: + for attempt in range(retries): + try: + r = client.chat.completions.create( + model=model, + messages=messages, + max_tokens=max_tokens, + extra_body={"reasoning": {"enabled": False}}, + ) + usage = {"input": r.usage.prompt_tokens, "output": r.usage.completion_tokens} + return (r.choices[0].message.content or ""), usage + except Exception: + if attempt == retries - 1: + raise + time.sleep(2 ** attempt) + raise RuntimeError("unreachable") diff --git a/tests/test_datagen.py b/tests/test_datagen.py new file mode 100644 index 0000000..f8cb11c --- /dev/null +++ b/tests/test_datagen.py @@ -0,0 +1,16 @@ +from buddhagpt.datagen import build_messages, parse_pairs, dedupe + +def test_build_messages_varies_templates(): + sutta = {"uid": "mn21", "title": "T", "text": "x" * 900} + prompts = {build_messages(sutta, v)[1]["content"] for v in range(6)} + assert len(prompts) == 6 # rotating templates, not one fixed prompt + +def test_parse_pairs_extracts_json_lines(): + out = '{"question": "Q1?", "answer": "A1"}\n{"question": "Q2?", "answer": "A2"}' + assert len(parse_pairs(out, uid="mn21")) == 2 + +def test_dedupe_drops_near_duplicates(): + pairs = [{"question": "What is craving?", "answer": "a"}, + {"question": "What is craving?", "answer": "b"}, + {"question": "How does one practice metta?", "answer": "c"}] + assert len(dedupe(pairs)) == 2 diff --git a/uv.lock b/uv.lock index 23aa09e..66550a4 100644 --- a/uv.lock +++ b/uv.lock @@ -60,6 +60,7 @@ dependencies = [ { name = "anthropic" }, { name = "lancedb" }, { name = "mlx-lm" }, + { name = "openai" }, { name = "pyyaml" }, { name = "sentence-transformers" }, ] @@ -74,6 +75,7 @@ requires-dist = [ { name = "anthropic", specifier = ">=0.122.0" }, { name = "lancedb", specifier = ">=0.37.1" }, { name = "mlx-lm", specifier = ">=0.31.3" }, + { name = "openai", specifier = ">=3.1.0" }, { name = "pyyaml", specifier = ">=6.0.3" }, { name = "sentence-transformers", specifier = ">=5.7.0" }, ] @@ -280,6 +282,19 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/7e/f5/f66802a942d491edb555dd61e3a9961140fd64c90bce1eafd741609d334d/httpcore-1.0.9-py3-none-any.whl", hash = "sha256:2d400746a40668fc9dec9810239072b40b4484b640a8c38fd654a024c7a1bf55", size = 78784, upload-time = "2025-04-24T22:06:20.566Z" }, ] +[[package]] +name = "httpcore2" +version = "2.10.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "h11" }, + { name = "truststore" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/a9/83/a896fc59940fc5a6e2aff3a4be1d92fa890112936803b331cae75a993c34/httpcore2-2.10.0.tar.gz", hash = "sha256:13c0cc3d1919d4f28457f60cd2c2abe04113a8af184ccf1142811beba936f9dc", size = 67427, upload-time = "2026-08-09T09:11:32.123Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/e5/4f/d149104195a35e2853a2fc203a8e3477747e58c80e17dda686dace174383/httpcore2-2.10.0-py3-none-any.whl", hash = "sha256:7df06cfb34070cae4f7c89be69dc1095eca138e9704ceffb98d25c1912ab6f01", size = 83000, upload-time = "2026-08-09T09:11:29.555Z" }, +] + [[package]] name = "httpx" version = "0.28.1" @@ -295,6 +310,32 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/2a/39/e50c7c3a983047577ee07d2a9e53faf5a69493943ec3f6a384bdc792deb2/httpx-0.28.1-py3-none-any.whl", hash = "sha256:d909fcccc110f8c7faf814ca82a9a4d816bc5a6dbfea25d6591d6985b8ba59ad", size = 73517, upload-time = "2024-12-06T15:37:21.509Z" }, ] +[[package]] +name = "httpx2" +version = "2.10.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "anyio", marker = "sys_platform != 'emscripten'" }, + { name = "httpcore2", marker = "sys_platform != 'emscripten'" }, + { name = "httpx2-jsfetch", marker = "sys_platform == 'emscripten'" }, + { name = "idna" }, + { name = "truststore", marker = "sys_platform != 'emscripten'" }, + { name = "typing-extensions", marker = "python_full_version < '3.13'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/bd/3d/f9a8c07a3884f3e5b26205e8436a18b3af61c5d53192c3bea235574dbbec/httpx2-2.10.0.tar.gz", hash = "sha256:8741d7329fe2c7885fc9ceb61c8217acfb87a85f75723714b89ebf7ad7196338", size = 98749, upload-time = "2026-08-09T09:11:33.24Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/b9/6d/a637d52449d98a6892d9a4dc0262587afdb6a66f201871842dce5a97b1c1/httpx2-2.10.0-py3-none-any.whl", hash = "sha256:5e3194a432701e1cc6f69a8b1b2fa199ef907013fede8d9a09a2c5b7b8141a18", size = 94355, upload-time = "2026-08-09T09:11:30.882Z" }, +] + +[[package]] +name = "httpx2-jsfetch" +version = "1.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/cd/c4/0e5636363151a2a1795e0a77617168b9ca438e1748ec05fc9b5687f93d64/httpx2_jsfetch-1.0.tar.gz", hash = "sha256:70a0e3eabfef7cce5ad9c629f7d01ca05e418f586646f4ddf14782e4c1454c60", size = 6872, upload-time = "2026-08-07T00:13:07.492Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/9b/43/832f631d32e4f1211caa2ba368317739fe71f0b8530e4c9d15dc454bac2a/httpx2_jsfetch-1.0-py3-none-any.whl", hash = "sha256:cb916b707601e69a07721aabc8f3f6659be3a6893bc1ff5c6f9e02241df2da32", size = 6382, upload-time = "2026-08-07T00:13:06.567Z" }, +] + [[package]] name = "huggingface-hub" version = "1.27.0" @@ -852,6 +893,25 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/a8/64/3708a90d1ebe202ffdeb7185f878a3c84d15c2b2c31858da2ce0583e2def/nvidia_nvtx-13.0.85-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:cb7780edb6b14107373c835bf8b72e7a178bac7367e23da7acb108f973f157a6", size = 148878, upload-time = "2025-09-04T08:28:53.627Z" }, ] +[[package]] +name = "openai" +version = "3.1.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "anyio" }, + { name = "distro" }, + { name = "httpx2" }, + { name = "jiter" }, + { name = "pydantic" }, + { name = "sniffio" }, + { name = "tqdm" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/58/9b/d45911bf9abfc5a754d800d79fe56e4dcf6e7b679d6ff4b7e9689b56bc02/openai-3.1.0.tar.gz", hash = "sha256:3ae7190da63f718727f9c525740d3f713e85553d2bf1d0cc9247346bd9063a4d", size = 1131016, upload-time = "2026-08-14T23:49:46.325Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/d3/45/84b9e909968c0f4c3112a828bba9aa8be6a5ee322fcb3d781145ee53d88e/openai-3.1.0-py3-none-any.whl", hash = "sha256:f06d5a5c8d06a0aead3a67d33fe5a3c3c31540a0f7a8017b02ba8cd99bc5c736", size = 1670762, upload-time = "2026-08-14T23:49:43.998Z" }, +] + [[package]] name = "packaging" version = "26.3" @@ -1541,6 +1601,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/f0/ac/229b7d4589d2e5937310e72c6d46e89599d16a4a12b479ffa1499fee8eb8/triton-3.7.1-cp314-cp314t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:10ba85fa2cca4a2fbdeb36bf1cb082f2c252bda55bf9fccd74f65ec5bc647e68", size = 197824404, upload-time = "2026-06-17T19:53:42.772Z" }, ] +[[package]] +name = "truststore" +version = "0.10.4" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/53/a3/1585216310e344e8102c22482f6060c7a6ea0322b63e026372e6dcefcfd6/truststore-0.10.4.tar.gz", hash = "sha256:9d91bd436463ad5e4ee4aba766628dd6cd7010cf3e2461756b3303710eebc301", size = 26169, upload-time = "2025-08-12T18:49:02.73Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/19/97/56608b2249fe206a67cd573bc93cd9896e1efb9e98bce9c163bcdc704b88/truststore-0.10.4-py3-none-any.whl", hash = "sha256:adaeaecf1cbb5f4de3b1959b42d41f6fab57b2b1666adb59e89cb0b53361d981", size = 18660, upload-time = "2025-08-12T18:49:01.46Z" }, +] + [[package]] name = "typer" version = "0.27.1"