"""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()