feat: synthetic instruction generation via OpenRouter deepseek-v4-flash
This commit is contained in:
@@ -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",
|
||||
]
|
||||
|
||||
82
scripts/gen_data.py
Normal file
82
scripts/gen_data.py
Normal file
@@ -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)
|
||||
69
src/buddhagpt/datagen.py
Normal file
69
src/buddhagpt/datagen.py
Normal file
@@ -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
|
||||
30
src/buddhagpt/llm.py
Normal file
30
src/buddhagpt/llm.py
Normal file
@@ -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")
|
||||
16
tests/test_datagen.py
Normal file
16
tests/test_datagen.py
Normal file
@@ -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
|
||||
69
uv.lock
generated
69
uv.lock
generated
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user