feat: CompassionBench bank + 4-system response collection
This commit is contained in:
147
scripts/collect_responses.py
Normal file
147
scripts/collect_responses.py
Normal file
@@ -0,0 +1,147 @@
|
||||
"""CompassionBench response collection: base, ft, ft_rag (local mlx) + kimi_k3 (frontier).
|
||||
|
||||
Deviates from the task-7 brief's collect_local() on purpose (controller ruling):
|
||||
the brief reloads the model via rag.answer() for every one of the 150 ft_rag
|
||||
prompts, which means 150 full 7B model loads. Instead we load the model +
|
||||
tokenizer ONCE per system pass, and for ft_rag inline retrieval (buddhagpt.index.search)
|
||||
+ prompt construction (buddhagpt.rag.build_prompt) around the single preloaded
|
||||
model/tokenizer.
|
||||
|
||||
Resumable: skips any (prompt_id, system) row already present in the output
|
||||
file, so a restart after an interruption does not redo work. Writes are
|
||||
flushed after every row so a kill -9 mid-run loses at most the in-flight
|
||||
generation.
|
||||
|
||||
Run order: kimi_k3 first (fast, validates the OpenRouter plumbing end to end
|
||||
before committing to the multi-hour local passes), then base, ft, ft_rag.
|
||||
|
||||
kimi_k3 uses a ThreadPoolExecutor(max_workers=8) (same pattern as
|
||||
scripts/gen_data.py's run(): a lock around the shared output-file append) since
|
||||
it's a remote API call and safely parallelizable. The three local mlx passes
|
||||
stay strictly sequential/single-stream — one GPU, no benefit to threading, and
|
||||
mlx generation is not thread-safe across concurrent calls on one model.
|
||||
"""
|
||||
import json
|
||||
import threading
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from pathlib import Path
|
||||
|
||||
from buddhagpt.collect import load_bench, load_done
|
||||
|
||||
BENCH = Path("eval/compassionbench.yaml")
|
||||
OUT = Path("data/responses.jsonl")
|
||||
BASE_MODEL = "mlx-community/Qwen2.5-7B-Instruct-4bit"
|
||||
FT_MODEL = "models/buddhagpt-7b-v1"
|
||||
RAG_DB = Path("data/lancedb")
|
||||
KIMI_MODEL = "moonshotai/kimi-k3"
|
||||
|
||||
|
||||
def run_frontier_pass(items: list[dict], out: Path, model: str = KIMI_MODEL, max_workers: int = 8):
|
||||
from buddhagpt.llm import openrouter_client, chat
|
||||
|
||||
client = openrouter_client()
|
||||
system = model.split("/")[-1].replace("-", "_")
|
||||
done = load_done(out)
|
||||
remaining = [it for it in items if (it["id"], system) not in done]
|
||||
print(f"[{system}] {len(remaining)}/{len(items)} remaining", flush=True)
|
||||
if not remaining:
|
||||
return
|
||||
|
||||
lock = threading.Lock()
|
||||
totals = {"input": 0, "output": 0}
|
||||
fail_count = 0
|
||||
done_count = 0
|
||||
out_f = out.open("a") # append: incremental persistence, survives interruption
|
||||
|
||||
def gen_one(it):
|
||||
nonlocal fail_count, done_count
|
||||
try:
|
||||
# kimi-k3 emits hidden reasoning tokens even with reasoning.enabled=False,
|
||||
# and they draw from the same max_tokens budget as the visible answer —
|
||||
# at 1000 tokens, harder prompts exhausted the budget mid-reasoning and
|
||||
# returned empty content (finish_reason="length", content=None). 3000
|
||||
# leaves enough room for reasoning + a full visible answer.
|
||||
text, usage = chat(client, model, [{"role": "user", "content": it["prompt"]}], max_tokens=3000)
|
||||
if not text or not text.strip():
|
||||
raise ValueError("empty response content (likely reasoning-token budget exhaustion)")
|
||||
except Exception as e:
|
||||
with lock:
|
||||
fail_count += 1
|
||||
print(f"[{system}] FAILED {it['id']}: {e}", flush=True)
|
||||
return
|
||||
with lock:
|
||||
totals["input"] += usage["input"]; totals["output"] += usage["output"]
|
||||
done_count += 1
|
||||
out_f.write(json.dumps({"prompt_id": it["id"], "system": system, "response": text}) + "\n")
|
||||
out_f.flush()
|
||||
print(f"[{system}] {done_count}/{len(remaining)} done ({it['id']})", flush=True)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=max_workers) as pool:
|
||||
list(pool.map(gen_one, remaining))
|
||||
|
||||
out_f.close()
|
||||
# Report raw token totals; actual $ cost is read from the OpenRouter activity
|
||||
# dashboard (per-token pricing for kimi-k3 not hardcoded here to avoid drift).
|
||||
print(
|
||||
f"[{system}] calls: {len(remaining)} | failed: {fail_count} | tokens: {totals}",
|
||||
flush=True,
|
||||
)
|
||||
|
||||
|
||||
def run_local_pass(items: list[dict], system: str, model_path: str, out: Path, use_rag: bool):
|
||||
from mlx_lm import load, generate
|
||||
|
||||
done = load_done(out)
|
||||
remaining = [it for it in items if (it["id"], system) not in done]
|
||||
print(f"[{system}] {len(remaining)}/{len(items)} remaining", flush=True)
|
||||
if not remaining:
|
||||
return
|
||||
|
||||
print(f"[{system}] loading model from {model_path} ...", flush=True)
|
||||
model, tokenizer = load(model_path)
|
||||
|
||||
if use_rag:
|
||||
from buddhagpt.index import search
|
||||
from buddhagpt.rag import build_prompt
|
||||
|
||||
with out.open("a") as f:
|
||||
for i, it in enumerate(remaining, 1):
|
||||
if use_rag:
|
||||
passages = search(RAG_DB, it["prompt"], k=4)
|
||||
messages = build_prompt(it["prompt"], passages)
|
||||
else:
|
||||
messages = [{"role": "user", "content": it["prompt"]}]
|
||||
prompt = tokenizer.apply_chat_template(
|
||||
messages, tokenize=False, add_generation_prompt=True
|
||||
)
|
||||
response = generate(model, tokenizer, prompt=prompt, max_tokens=500)
|
||||
f.write(json.dumps({"prompt_id": it["id"], "system": system, "response": response}) + "\n")
|
||||
f.flush()
|
||||
print(f"[{system}] {i}/{len(remaining)} done ({it['id']})", flush=True)
|
||||
|
||||
|
||||
def main():
|
||||
items = load_bench(BENCH)
|
||||
OUT.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# 1. Frontier reference first — fast, validates plumbing before the long local runs.
|
||||
# Parallel (8 workers): remote API calls, safely concurrent.
|
||||
print("=== kimi_k3 (frontier reference) ===", flush=True)
|
||||
run_frontier_pass(items, OUT)
|
||||
|
||||
# 2. Local passes — each loads the model exactly once.
|
||||
print("=== base ===", flush=True)
|
||||
run_local_pass(items, "base", BASE_MODEL, OUT, use_rag=False)
|
||||
|
||||
print("=== ft ===", flush=True)
|
||||
run_local_pass(items, "ft", FT_MODEL, OUT, use_rag=False)
|
||||
|
||||
print("=== ft_rag ===", flush=True)
|
||||
run_local_pass(items, "ft_rag", FT_MODEL, OUT, use_rag=True)
|
||||
|
||||
total = sum(1 for _ in OUT.open())
|
||||
print(f"done. total rows: {total}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user